{"record":{"id":"be57a02397a81ac4","repo":"jax-ml/jax","slug":"argument-arg-of-type-type-arg-is-not-a-vali","errorCode":null,"errorMessage":"Argument '{arg}' of type {type(arg)} is not a valid JAX type.","messagePattern":"Argument '(.+?)' of type (.+?) is not a valid JAX type\\.","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/dispatch.py","lineNumber":287,"sourceCode":"    elif eqn.primitive is shard_map.shard_map_p:\n      mesh = eqn.params['mesh']\n      if isinstance(mesh, AbstractMesh):\n        continue\n      source_info = SourceInfo(eqn.source_info, eqn.primitive.name)\n      out.extend((NamedSharding(mesh, spec), source_info)\n                 for spec in [*eqn.params['in_specs'], *eqn.params['out_specs']])\n    elif eqn.primitive is device_put_p:\n      source_info = SourceInfo(eqn.source_info, eqn.primitive.name)\n      out.extend((s, source_info) for s in eqn.params['devices']\n                 if isinstance(s, Sharding) and s.memory_kind is not None)\n  for subjaxpr in core.subjaxprs(jaxpr):\n    out.extend(get_intermediate_shardings(subjaxpr))\n  return out\n\n\ndef check_arg(arg: Any):\n  if not core.valid_jaxtype(arg):\n    raise TypeError(f\"Argument '{arg}' of type {type(arg)} is not a valid \"\n                    \"JAX type.\")\n\n\ndef needs_check_special() -> bool:\n  return config.debug_infs.value or config.debug_nans.value\n\ndef check_special(name: str, bufs: Sequence[basearray.Array]) -> None:\n  if needs_check_special():\n    for buf in bufs:\n      _check_special(name, buf.dtype, buf)\n\n\ndef check_special_array(name: str, arr: array.ArrayImpl) -> array.ArrayImpl:\n  if needs_check_special():\n    if dtypes.issubdtype(arr.dtype, np.inexact):\n      for buf in arr._arrays:\n        _check_special(name, buf.dtype, buf)\n  return arr","sourceCodeStart":269,"sourceCodeEnd":305,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/dispatch.py#L269-L305","documentation":"JAX validates every argument passed into a transform or dispatch path with core.valid_jaxtype; anything that is not a JAX-compatible type (array, tracer, standard scalar container) is rejected with this TypeError. It is raised by check_arg, which guards entry points like grad/jacfwd/pjit dispatch. Its purpose is to fail fast before tracing so users get a clear message instead of a cryptic tracer error.","triggerScenarios":"Passing a non-JAX object to jax.grad, jax.jacfwd, jit-compiled/pjit functions, or other transforms: e.g. a numpy object-dtype array, a Python custom class, a torch tensor, None, or a string used as a leaf in a pytree.","commonSituations":"Passing PyTorch tensors or raw Python objects into JAX functions; object-dtype numpy arrays from pandas; dictionaries with non-array leaves; strings/None accidentally used as data.","solutions":["Convert the offending argument with jnp.asarray(x) (and a valid dtype) before passing it in","If using foreign tensors, convert via numpy: jnp.asarray(x.numpy()) or jnp.from_dlpack(x)","Check for object dtype: np.asarray(x).dtype == object and fix the data source","Inspect the traceback to identify which argument named in the message is invalid"],"exampleFix":"// before\nloss = jax.grad(model)(raw_python_list, params)\n// after\nloss = jax.grad(model)(jnp.asarray(raw_python_list, dtype=jnp.float32), params)","handlingStrategy":"type-guard","validationCode":"import jax\ndef is_jax_type(x):\n    try:\n        jax.core.typeof(x); return True\n    except TypeError:\n        return False\nvals = [v for v in args if not is_jax_type(v)]\nassert not vals, f'non-JAX args: {vals}'","typeGuard":"def is_jax_arg(x) -> bool:\n    import jax.numpy as jnp\n    return hasattr(x, 'dtype') and hasattr(x, 'shape') or isinstance(x, (int, float, complex, bool))","tryCatchPattern":"try:\n    result = jax.grad(f)(*args)\nexcept TypeError as e:\n    if 'not a valid JAX type' in str(e):\n        args = tuple(jnp.asarray(a) if not hasattr(a, 'dtype') else a for a in args)","preventionTips":["Always convert inputs with jnp.asarray before passing to JAX transforms","Validate pytree leaves with jax.tree_util.tree_map(lambda l: jnp.asarray(l), data)","Avoid object-dtype numpy arrays; cast with astype(float)"],"tags":["jax","type-validation","argument-validation"],"backgroundTag":"invalid-argument-type","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}