{"record":{"id":"802c855eedac77d5","repo":"jax-ml/jax","slug":"name-wrapped-function-must-be-passed-at-least-on","errorCode":null,"errorMessage":"{name} wrapped function must be passed at least one argument containing an array or axis_size must be specified, got empty *args={args} and **kwargs={kwargs}","messagePattern":"(.+?) wrapped function must be passed at least one argument containing an array or axis_size must be specified, got empty \\*args=(.+?) and \\*\\*kwargs=(.+?)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/api.py","lineNumber":1301,"sourceCode":"def _check_ema_unmapped_args(ema, args_flat, in_axes_flat):\n  if ema is None:\n    return\n  for a, i in zip(args_flat, in_axes_flat):\n    if i is None:\n      aval = core.typeof(a)\n      spec = set(sharding_impls.flatten_spec(aval.sharding.spec))\n      if any(e in spec for e in ema):\n        raise ValueError(\n            \"Unmapped values passed to vmap cannot be sharded along the mesh\"\n            f\" axis you are vmapping over. Got type: {aval.str_short(True)},\"\n            f\" in_axes: {i} and vmapped mesh axis: {ema}\")\n\ndef _mapped_axis_size(fn, tree, vals, dims, name, axis_size=None):\n  if not vals:\n    if axis_size is not None:\n      return axis_size\n    args, kwargs = tree_unflatten(tree, vals)\n    raise ValueError(\n        f\"{name} wrapped function must be passed at least one argument \"\n        \"containing an array or axis_size must be specified, got empty \"\n        f\"*args={args} and **kwargs={kwargs}\"\n    )\n\n  def _get_axis_size(name: str, x, axis: int) -> core.AxisSize | None:\n    shape: tuple[core.AxisSize, ...] = ()\n    try:\n      shape = np.shape(x)\n      return shape[axis]\n    except (IndexError, TypeError) as e:\n      if not core.valid_jaxtype(x) or not isinstance(axis, int):\n        return None  # Suppress the check for custom vmappable types.\n      if core.typeof(x).is_high:\n        raise ValueError(\n            f\"{name} was requested to map a value of non-array type \"\n            f\"{core.typeof(x)} along axis {axis}, but non-array types can't \"\n            \"be mapped along an integer axis. Instead pass a mapping spec (a \"","sourceCodeStart":1283,"sourceCodeEnd":1319,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/api.py#L1283-L1319","documentation":"Raised by vmap/pmap-style transforms when the wrapped function receives no arguments containing arrays (or any mapped leaves) and no axis_size was given, so the mapped axis size cannot be inferred.","triggerScenarios":"jax.vjp-like call such as jax.vmap(lambda: ... )() or jax.vmap(f)() where all arguments are empty/pytrees with no leaves, or jax.linearize/jax.vjp on a zero-argument function, without passing axis_size.","commonSituations":"Refactoring a function to take configuration via closure instead of arguments; calling vmap over a function of only Python scalars/static values; testing with dummy empty inputs.","solutions":["Pass axis_size=N to vmap so the mapped size is explicit","Add at least one array argument to the function and map it via in_axes","Move constants into the function's arguments as arrays"],"exampleFix":"// before\njax.vmap(lambda: jnp.arange(3) + 1)()\n// after\njax.vmap(lambda n: jnp.arange(3) + n, axis_size=8)()\n# or\njax.vmap(lambda n: jnp.arange(3) + n)(jnp.ones(8))","handlingStrategy":"validation","validationCode":"if not any(hasattr(l, 'ndim') for l in tree_leaves((args, kwargs))):\n    assert axis_size is not None, 'pass axis_size when no array args are mapped'","typeGuard":null,"tryCatchPattern":null,"preventionTips":["Always pass axis_size for vmap over zero-array functions","Prefer array arguments over closures for mapped values","Wrap risky calls: jax.vmap(f, axis_size=N) by default in such utilities"],"tags":["jax","vmap","axis-size","api-misuse"],"backgroundTag":"missing-required-argument","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}