{"record":{"id":"26deec772604e08a","repo":"jax-ml/jax","slug":"ragged-dot-general-requires-that-group-sizes-dtype","errorCode":null,"errorMessage":"ragged_dot_general requires that group_sizes.dtype is subtype of np.integer.","messagePattern":"ragged_dot_general requires that group_sizes\\.dtype is subtype of np\\.integer\\.","errorType":"validation","errorClass":"TypeError","httpStatus":null,"severity":"error","filePath":"jax/_src/lax/lax.py","lineNumber":6467,"sourceCode":"  )\n  if mode == RaggedDotMode.RAGGED_CONTRACTING:\n    out_shape = (num_groups,) + out_shape\n  return out_shape\n\n\ndef _ragged_dot_general_dtype_rule(\n    lhs: Array,\n    rhs: Array,\n    group_sizes: Array,\n    *,\n    ragged_dot_dimension_numbers: RaggedDotDimensionNumbers,\n    precision,\n    preferred_element_type: DTypeLike | None,\n    group_offset,\n    out_sharding,\n) -> np.dtype:\n  if not dtypes.issubdtype(group_sizes.dtype, np.integer):\n    raise TypeError(\n        'ragged_dot_general requires that '\n        'group_sizes.dtype is subtype of np.integer.'\n    )\n  # defer the output dtype to dot_general, which is part of the _ragged_dot_general_impl.\n  return _dot_general_dtype_rule(\n      lhs,\n      rhs,\n      dimension_numbers=ragged_dot_dimension_numbers.dot_dimension_numbers,\n      precision=precision,\n      preferred_element_type=preferred_element_type,\n      out_sharding=None,\n      name='lax.ragged_dot_general',\n  )\n\n\ndef _ragged_dot_general_jvp_rule(\n    primals, tangents, ragged_dot_dimension_numbers,\n    precision, preferred_element_type, group_offset, out_sharding","sourceCodeStart":6449,"sourceCodeEnd":6485,"githubUrl":"https://github.com/jax-ml/jax/blob/1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb/jax/_src/lax/lax.py#L6449-L6485","documentation":"The dtype rule for ragged_dot_general requires group_sizes to be an integer array. Float group sizes (e.g. from softmax weights or averaged counts) are rejected because group sizes drive indexing/cumsum arithmetic.","triggerScenarios":"Passing a group_sizes array with a floating dtype (e.g. result of jnp.mean, jnp.float32 counts, or a tracer carrying float dtype) to ragged_dot_general.","commonSituations":"Computing token-per-expert counts with float ops (softmax, division) and forgetting to cast; loading counts from a float numpy array or CSV.","solutions":["Cast group_sizes to an integer dtype: group_sizes.astype(jnp.int32)","Produce counts with integer ops from the start (e.g. jnp.sum of one-hot routing, bincount)","Validate dtypes before the call (see type guard below)"],"exampleFix":"// before\ngroup_sizes = counts.mean(axis=0)          # float32\nout = ragged_dot_general(x, w, group_sizes, dn, mode=mode)\n// after\ngroup_sizes = counts.sum(axis=0).astype(jnp.int32)\nout = ragged_dot_general(x, w, group_sizes, dn, mode=mode)","handlingStrategy":"type-guard","validationCode":"group_sizes = jnp.asarray(group_sizes)\nif not jnp.issubdtype(group_sizes.dtype, jnp.integer):\n    group_sizes = group_sizes.astype(jnp.int32)","typeGuard":"def is_int_array(a) -> bool:\n    return dtypes.issubdtype(np.asarray(a).dtype, np.integer)","tryCatchPattern":null,"preventionTips":["Always .astype(jnp.int32) counts derived from float math","Prefer bincount/sum for counts"],"tags":["jax","ragged-dot","dtype-validation"],"backgroundTag":"dtype-validation-failed","analyzedSha":"1e1c6a8fc06dfcd1247076ec5cae4640cea5d7bb","analyzedAt":"2026-08-27T09:53:25.647Z","schemaVersion":2},"datasetVersion":"2026-08-27T13:17:12.746Z"}