{"record":{"id":"f57d23ef208b9280","repo":"jax-ml/jax","slug":"input-type-mismatch-for-prim","errorCode":null,"errorMessage":"input type mismatch for {_prim}","messagePattern":"input type mismatch for (.+?)","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/hijax.py","lineNumber":371,"sourceCode":"  if isinstance(ct, ad_util.Zero):\n    return ad_util.Zero(core.unmapped_aval(axis_data.size, d, ct.aval,\n                                           axis_data.explicit_mesh_axis))\n  return ct\n\n\ncall_hi_primitive_p = core.Primitive(\"call_hi_primitive\")\ncall_hi_primitive_p.multiple_results = True\ncall_hi_primitive_p.skip_canonicalization = True\ncall_hi_primitive_p.is_high = lambda *args, _prim: True\ncall_hi_primitive_p.is_effectful = lambda params: bool(params['_prim'].effects)\n@call_hi_primitive_p.def_effectful_abstract_eval\ndef _call_hi_primitive_abstract_eval(*_args, _prim):\n  return _prim.out_avals_flat, _prim.effects\n\ndef _call_hi_primitive_typecheck(_ctx_factory, *in_atoms_flat, _prim):\n  in_avals = [x.aval for x in in_atoms_flat]\n  if not all(map(core.typematch, in_avals, _prim.in_avals_flat)):\n    raise TypeError(f\"input type mismatch for {_prim}\")\n  _prim.check()\n  return _prim.out_avals_flat, _prim.effects\ncore.custom_typechecks[call_hi_primitive_p] = _call_hi_primitive_typecheck\n\ndef _call_hi_primitive_staging(trace, source_info, *args_flat, _prim):\n  trace.frame.is_high = True\n  args = tree_unflatten(_prim.in_tree, args_flat)\n  ans = _prim.staging(trace, source_info, *args)\n  return tree_leaves_checked(_prim.out_tree, ans)\npe.custom_staging_rules[call_hi_primitive_p] = _call_hi_primitive_staging\n\ndef _call_hi_primitive_to_lojax(*args_flat, _prim):\n  args = tree_unflatten(_prim.in_tree, args_flat)\n  ans = _prim.expand(*args)\n  return tree_leaves_checked(_prim.out_tree, ans)\ncall_hi_primitive_p.to_lojax = _call_hi_primitive_to_lojax\n\ndef _call_hi_primitive_prettyprint(eqn, context, settings):","sourceCodeStart":353,"sourceCodeEnd":389,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/hijax.py#L353-L389","documentation":"A HiPrim application is being staged/traced with input avals that don't typematch the avals recorded when the primitive was created. The typecheck registered for call_hi_primitive_p compares each input aval with _prim.in_avals_flat and raises TypeError on mismatch.","triggerScenarios":"Re-binding a traced HiPrim (e.g. re-invoking a stored traced primitive, or cacheing a primitive across calls) with arguments of different dtype/shape/weak_type than the original trace inputs.","commonSituations":"Reusing a cached/traced primitive with a differently-typed input (e.g. int vs float, or shape change after padding), or hitting stale closures after dtype promotion.","solutions":["Re-create the primitive (re-trace) for the new input types instead of reusing the old instance","Ensure input dtypes/shapes are consistent, e.g. cast with jnp.asarray(x, dtype=...) before the call","Check for accidental integer-literal or weak-typed inputs (Python scalars) vs arrays"],"exampleFix":"# before\np = MyPrim(args1)  # traced with float32\np(args2_int)\n# after\np = MyPrim(jnp.asarray(args2_int, args1.dtype))  # retrace with correct types","handlingStrategy":"type-guard","validationCode":"import jax.numpy as jnp\nargs = [jnp.asarray(a, dtype=p.in_avals_flat[i].dtype) for i, a in enumerate(args_flat)]\nassert all(core.typematch(a.aval, b) for a, b in zip(traced_args, p.in_avals_flat))","typeGuard":"def inputs_typematch(prim, args_flat) -> bool:\n    import jax._src.core as core\n    return all(map(core.typematch,\n                   [getattr(a, 'aval', a) for a in args_flat],\n                   prim.in_avals_flat))","tryCatchPattern":"try:\n    p(*args)\nexcept TypeError as e:\n    if 'input type mismatch' in str(e):\n        p = type(p)(*[jnp.asarray(a, p.in_avals_flat[i].dtype)\n                      for i, a in enumerate(args)])\n        return p(*args)\n    raise","preventionTips":["Normalize input dtypes/shapes before invoking traced primitives","Re-trace primitives whenever input types change","Avoid mixing Python scalars with typed arrays at the boundary"],"tags":["jax","type-mismatch","tracing","custom-primitive"],"backgroundTag":"traced-input-type-mismatch","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}