{"record":{"id":"d68e8b8a008f90ad","repo":"jax-ml/jax","slug":"type-prim-name-transpose-should-return-non","errorCode":null,"errorMessage":"{type(_prim).__name__}.transpose should return None or a dict of backward-pass log entries, got {type(log).__name__}","messagePattern":"(.+?)\\.transpose should return None or a dict of backward-pass log entries, got (.+?)","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/hijax.py","lineNumber":520,"sourceCode":"    del params['has_sres']\n  return core._pp_eqn(eqn.replace(params=params), context, settings)\ncore.pp_eqn_rules[call_hi_primitive_linearized_p] = _call_hi_primitive_linearized_prettyprint\n\ndef _call_hi_primitive_jvp(primals, tangents, *, _prim):\n  primals = tree_unflatten(_prim.in_tree, primals)\n  tangents = tree_unflatten(_prim.in_tree, tangents)\n  out_primals, out_tangents = _prim.jvp(primals, tangents)\n  out_primals_flat = tree_leaves_checked(_prim.out_tree, out_primals)\n  out_tangents_flat = _prim.out_tree.flatten_up_to(out_tangents)\n  return out_primals_flat, out_tangents_flat\nad.primitive_jvps[call_hi_primitive_p] = _call_hi_primitive_jvp\n\ndef _call_hi_primitive_transpose(cts_flat, *primals_flat, _prim):\n  cts = tree_unflatten(_prim.out_tree, cts_flat)\n  primals = tree_unflatten(_prim.in_tree, primals_flat)\n  log = _prim.transpose(cts, *primals)  # a returned dict logs entries\n  if log is not None and type(log) is not dict:\n    raise TypeError(\n        f\"{type(_prim).__name__}.transpose should return None or a dict of \"\n        f\"backward-pass log entries, got {type(log).__name__}\")\n  return log\nad.fancy_transposes[call_hi_primitive_p] = _call_hi_primitive_transpose\n\ndef _call_hi_primitive_dce(used_outs_flat, eqn):\n  _prim = eqn.params['_prim']\n  used_out = tree_unflatten(_prim.out_tree, used_outs_flat)\n  used_ins, produced_outs, new_prim = _prim.dce(used_out)\n  if new_prim is None:\n    return [False] * len(eqn.invars), None\n  name = f'{type(_prim).__name__}.dce'\n  used_ins_flat = api.tuptree_flags(\n      used_ins, _prim.in_tree, 'used_ins',\n      f'the first (used inputs) return value of {name}')\n  produced_outs_flat = api.tuptree_flags(\n      produced_outs, _prim.out_tree, 'produced_outs',\n      f'the second (produced outputs) return value of {name}')","sourceCodeStart":502,"sourceCodeEnd":538,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/hijax.py#L502-L538","documentation":"The fancy-transpose handler for call_hi_primitive_p calls _prim.transpose(cts, *primals); its return must be None or a dict of log entries. The actual transposed values are delivered through the machinery, so returning cotangents directly is a TypeError.","triggerScenarios":"A custom HiPrim transpose implementation returns cotangents (tuple/array) instead of None or a logs dict.","commonSituations":"Writing transpose with the conventional JAX convention (return input cotangents) rather than hijax's logging convention.","solutions":["Rework transpose to pass results via its output-accumulator mechanism and return None or a logs dict","Follow the HiPrim.transpose docstring/examples for the expected signature `transpose(self, out_ct, *maybe_accums)`"],"exampleFix":"# before\ndef transpose(self, out_ct, *args):\n    return args\n# after\ndef transpose(self, out_ct, *args):\n    ...\n    return None","handlingStrategy":"try-catch","validationCode":null,"typeGuard":"def transpose_returns_valid(p) -> bool:\n    out = p.transpose(dummy_ct, *dummy_primals)\n    return out is None or type(out) is dict","tryCatchPattern":"try:\n    jax.linear_transpose(f, x)(y)\nexcept TypeError as e:\n    if 'transpose should return' in str(e):\n        raise RuntimeError('transpose must return None or a logs dict') from e\n    raise","preventionTips":["Treat transpose's return as logging-only","Follow hijax examples when writing transpose rules"],"tags":["jax","autodiff","transpose","api-contract"],"backgroundTag":"autodiff-rule-return-type","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}