{"record":{"id":"dedb20ce2d112365","repo":"jax-ml/jax","slug":"unexpected-type-ty-for-value-v","errorCode":null,"errorMessage":"Unexpected type {ty} for value {v}","messagePattern":"Unexpected type (.+?) for value (.+?)","errorType":"validation","errorClass":"ValueError","httpStatus":null,"severity":"error","filePath":"jax/_src/pallas/mosaic_gpu/primitives.py","lineNumber":3544,"sourceCode":"    case (ShapeDtypeStruct(), mgpu.FragmentedArray()):\n      mlir_dtype = mgpu_utils.dtype_to_ir_type(ty.dtype)\n      if v.mlir_dtype != mlir_dtype:\n        raise ValueError(\n            f\"Array dtype mismatch: expected {v.mlir_dtype} got {mlir_dtype}.\"\n        )\n      if ty.shape != v.shape:\n        raise ValueError(\n            f\"Array shape mismatch: expected {ty.shape} got {v.shape}.\"\n        )\n      if v.layout != ty.layout.to_mgpu():\n        raise ValueError(\n            f\"Array layout mismatch: expected {v.layout} got {ty.layout.to_mgpu()}.\"\n        )\n    case (SomeLayout(), mgpu.FragmentedArray()):\n      if ty.to_mgpu() != v.layout:\n        raise ValueError(f\"Unexpected layout for {v} (expected: {ty})\")\n    case _:\n      raise ValueError(f\"Unexpected type {ty} for value {v}\")\n\n\ndef _inline_mgpu_flat_transformed_args(\n    ctx: lowering.LoweringRuleContext,\n    flat_args_and_transforms,\n    flat_arg_types,\n    pytree_args,\n    pytree_ref_transforms,\n  ) -> Sequence[ir.Value | mgpu.FragmentedArray]:\n  flat_args = flat_args_and_transforms[:pytree_args.num_leaves]\n  flat_arg_avals = ctx.avals_in[:pytree_args.num_leaves]\n  ref_transforms = pytree_ref_transforms.unflatten(flat_args_and_transforms[pytree_args.num_leaves:])\n  ref_transform_avals = pytree_ref_transforms.unflatten(ctx.avals_in[pytree_args.num_leaves:])\n  is_wg_semantics = (\n      ctx.module_ctx.lowering_semantics == mgpu.LoweringSemantics.Warpgroup\n  )\n  is_warp_semantics = (\n      ctx.module_ctx.primitive_semantics == gpu_core.PrimitiveSemantics.Warp","sourceCodeStart":3526,"sourceCodeEnd":3562,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/pallas/mosaic_gpu/primitives.py#L3526-L3562","documentation":"The inline_mgpu lane-semantics type checker only understands (RefType, MemRef value), (ShapeDtypeStruct, FragmentedArray) and (SomeLayout, FragmentedArray) pairs. Any other (type, value) combination falls through to this generic 'Unexpected type' error.","triggerScenarios":"Declaring a type entry that is neither RefType nor ShapeDtypeStruct nor SomeLayout-derived, or passing a runtime value that is not an ir.Value memref or FragmentedArray (e.g. a python scalar, bool).","commonSituations":"Passing predicates as plain python bools instead of FragmentedArray splats; custom type annotations in arg_types.","solutions":["Convert the value to a supported runtime representation (splat scalars into FragmentedArray via mgpu.splat)","Use only RefType / ShapeDtypeStruct / layout types in arg_types and return_type"],"exampleFix":"# before\nf(True)\n# after\nf(mgpu.splat(True, layout))  # or use mgpu.c for MLIR constants inside the impl","handlingStrategy":"type-guard","validationCode":"def supported_pair(ty, v):\n    return (isinstance(ty, RefType) and isinstance(v.type, ir.MemRefType)) or (isinstance(ty, (ShapeDtypeStruct, SomeLayout)) and isinstance(v, mgpu.FragmentedArray))\nassert supported_pair(ty, v)","typeGuard":"def supported_pair(ty, v):\n    return (isinstance(ty, RefType) and isinstance(v.type, ir.MemRefType)) or (isinstance(ty, (ShapeDtypeStruct, SomeLayout)) and isinstance(v, mgpu.FragmentedArray))","tryCatchPattern":null,"preventionTips":["Convert scalars/bools to FragmentedArray splats before passing","Restrict annotations to the three supported type families"],"tags":["jax","pallas","inline-mgpu","type-validation"],"backgroundTag":"unsupported-argument-type","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}