{"record":{"id":"7af541d7ad10cfd7","repo":"jax-ml/jax","slug":"devices-argument-to-pmap-must-be-non-empty-or-n","errorCode":null,"errorMessage":"'devices' argument to pmap must be non-empty, or None.","messagePattern":"'devices' argument to pmap must be non-empty, or None\\.","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pmap.py","lineNumber":280,"sourceCode":"  which runs on the first six devices and one on the remaining two:\n\n  >>> from functools import partial\n  >>> @partial(pmap, axis_name='i', devices=jax.devices()[:6])\n  ... def f1(x):\n  ...   return x / jax.lax.psum(x, axis_name='i')\n  >>>\n  >>> @partial(pmap, axis_name='i', devices=jax.devices()[-2:])\n  ... def f2(x):\n  ...   return jax.lax.psum(x ** 2, axis_name='i')\n  >>>\n  >>> print(f1(jnp.arange(6.)))  # doctest: +SKIP\n  [0.         0.06666667 0.13333333 0.2        0.26666667 0.33333333]\n  >>> print(f2(jnp.array([2., 3.])))  # doctest: +SKIP\n  [ 13.  13.]\n  \"\"\"\n  if devices is not None:\n    if not devices:\n      raise ValueError(\"'devices' argument to pmap must be non-empty, or None.\")\n    devices = tuple(devices)\n  axis_name, static_broadcasted_tuple, donate_tuple = _prepare_pmap(\n      fun, axis_name, static_broadcasted_argnums, donate_argnums, in_axes,\n      out_axes)\n  wrapped_fun = _pmap_wrap_init(fun, static_broadcasted_tuple)\n  out_axes_flat, out_axes_tree = tree_flatten(out_axes)\n  out_axes_flat = tuple(out_axes_flat)\n\n  def infer_params(*args, **kwargs):\n    process_count = xb.process_count(backend)\n    trace_state_clean = core.trace_state_clean()\n    dyn_f, dyn_argnums, dyn_args = _get_dyn_args(\n        wrapped_fun, static_broadcasted_tuple, args)\n    dyn_args_flat, dyn_args_tree = tree_flatten((dyn_args, kwargs))\n    in_axes_flat = _get_in_axes_flat(\n        in_axes, dyn_argnums, dyn_args, kwargs, len(dyn_args_flat),\n        dyn_args_tree)\n    local_axis_size = _mapped_axis_size(dyn_args_flat, in_axes_flat)","sourceCodeStart":262,"sourceCodeEnd":298,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pmap.py#L262-L298","documentation":"`jax.pmap`'s `devices` parameter must be None (use all devices) or a non-empty sequence. An empty list/tuple is rejected because pmap needs at least one device to map over.","triggerScenarios":"Calling `jax.pmap(f, devices=[])` or passing an empty device list computed dynamically, e.g. `devices=jax.devices('gpu')` when no GPU backend is visible.","commonSituations":"Selecting devices by backend string that returns nothing (no GPU/TPU visible); filtering devices and getting an empty result; CI environments without accelerators.","solutions":["Check the device list is non-empty before passing it (fall back to None)","Fix backend visibility: install correct CUDA/JAX version so `jax.devices('gpu')` returns devices","Pass `devices=None` to use all available devices"],"exampleFix":"# before\nf = jax.pmap(fn, devices=jax.devices('gpu'))  # empty if no GPU\n# after\ndevs = jax.devices('gpu') or None\nf = jax.pmap(fn, devices=devs)","handlingStrategy":"validation","validationCode":"import jax\ndevs = jax.devices('gpu')\nif devs:\n    f = jax.pmap(fn, devices=devs)\nelse:\n    f = jax.jit(fn)  # fallback","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Never pass an empty devices list; use None","Check backend availability before selecting devices"],"tags":["jax","pmap","devices","empty-argument"],"backgroundTag":"empty-device-list","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}