{"record":{"id":"e30114756e033c12","repo":"jax-ml/jax","slug":"expected-unreduced-kind-to-be-of-type-jax-shardin","errorCode":null,"errorMessage":"Expected unreduced_kind to be of type `jax.sharding.UnreducedKind` but got {type(unreduced_kind)}","messagePattern":"Expected unreduced_kind to be of type `jax\\.sharding\\.UnreducedKind` but got (.+?)","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/core.py","lineNumber":2355,"sourceCode":"\ndef get_memory_space(memory_space):\n  assert memory_space is not None\n  return memory_space\n\n\ndef _check_mat(varying, unreduced, reduced, unreduced_kind):\n  if varying & unreduced:\n    raise ValueError(\n        \"varying and unreduced cannot have common mesh axes. Got\"\n        f\" varying={varying} and unreduced={unreduced}\")\n  if varying & reduced:\n    raise ValueError(\n        \"varying and reduced cannot have common mesh axes. Got\"\n        f\" varying={varying} and reduced={reduced}\")\n  assert not (varying & unreduced & reduced)\n\n  if unreduced_kind is not None and not isinstance(unreduced_kind, UnreducedKind):\n    raise TypeError(\n        \"Expected unreduced_kind to be of type `jax.sharding.UnreducedKind`\"\n        f\" but got {type(unreduced_kind)}\")\n  if not unreduced and unreduced_kind is not None:\n    raise ValueError(\n        \"`unreduced_kind` should be `None` when `unreduced` is an empty set.\"\n        f\" Got {unreduced_kind=} and {unreduced=}\")\n\ndef _canonicalize_mat(name, val):\n  if not isinstance(val, frozenset):\n    if not isinstance(val, set):\n      raise TypeError(\n          f\"{name} argument of ManualAxisType should \"\n          f\"of type `frozenset` or `set`. Got type {type(val)}\")\n    val = frozenset(val)\n  return val\n\n\n@immutable","sourceCodeStart":2337,"sourceCodeEnd":2373,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/core.py#L2337-L2373","documentation":"ManualAxisType's optional unreduced_kind annotation must be a jax.sharding.UnreducedKind (e.g. UnreducedKind.sum/min/max) or None. Passing a string, int, or other object is a type error.","triggerScenarios":"ManualAxisType(..., unreduced={'x'}, unreduced_kind='sum') or unreduced_kind=0 instead of the enum-like UnreducedKind instance.","commonSituations":"Assuming unreduced_kind is a free-form string; serializing/deserializing mats and losing the type; old JAX versions predating the typed enum.","solutions":["Pass jax.sharding.UnreducedKind.sum (or .min/.max) or None","On deserialization, map stored strings back to UnreducedKind members","Leave it None when there are no unreduced axes"],"exampleFix":"// before\nmat = ManualAxisType(unreduced={'x'}, unreduced_kind='sum')\n\n// after\nfrom jax.sharding import UnreducedKind\nmat = ManualAxisType(unreduced={'x'}, unreduced_kind=UnreducedKind.sum)","handlingStrategy":"type-guard","validationCode":"from jax.sharding import UnreducedKind\nassert unreduced_kind is None or isinstance(unreduced_kind, UnreducedKind)","typeGuard":"def valid_kind(k): return k is None or isinstance(k, UnreducedKind)","tryCatchPattern":null,"preventionTips":["Never pass strings for unreduced_kind","Map serialized strings back to UnreducedKind members on load"],"tags":["jax","sharding","unreduced-kind","type-error"],"backgroundTag":"invalid-enum-argument","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}