jax-ml/jax · error · TypeError

ragged_dot_general requires that group_sizes.dtype is subtyp

Error message

ragged_dot_general requires that group_sizes.dtype is subtype of np.integer.

What it means

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.

Source

Thrown at jax/_src/lax/lax.py:6467

  )
  if mode == RaggedDotMode.RAGGED_CONTRACTING:
    out_shape = (num_groups,) + out_shape
  return out_shape


def _ragged_dot_general_dtype_rule(
    lhs: Array,
    rhs: Array,
    group_sizes: Array,
    *,
    ragged_dot_dimension_numbers: RaggedDotDimensionNumbers,
    precision,
    preferred_element_type: DTypeLike | None,
    group_offset,
    out_sharding,
) -> np.dtype:
  if not dtypes.issubdtype(group_sizes.dtype, np.integer):
    raise TypeError(
        'ragged_dot_general requires that '
        'group_sizes.dtype is subtype of np.integer.'
    )
  # defer the output dtype to dot_general, which is part of the _ragged_dot_general_impl.
  return _dot_general_dtype_rule(
      lhs,
      rhs,
      dimension_numbers=ragged_dot_dimension_numbers.dot_dimension_numbers,
      precision=precision,
      preferred_element_type=preferred_element_type,
      out_sharding=None,
      name='lax.ragged_dot_general',
  )


def _ragged_dot_general_jvp_rule(
    primals, tangents, ragged_dot_dimension_numbers,
    precision, preferred_element_type, group_offset, out_sharding

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Cast group_sizes to an integer dtype: group_sizes.astype(jnp.int32)
  2. Produce counts with integer ops from the start (e.g. jnp.sum of one-hot routing, bincount)
  3. Validate dtypes before the call (see type guard below)

Example fix

// before
group_sizes = counts.mean(axis=0)          # float32
out = ragged_dot_general(x, w, group_sizes, dn, mode=mode)
// after
group_sizes = counts.sum(axis=0).astype(jnp.int32)
out = ragged_dot_general(x, w, group_sizes, dn, mode=mode)
Defensive patterns

Strategy: type-guard

Validate before calling

group_sizes = jnp.asarray(group_sizes)
if not jnp.issubdtype(group_sizes.dtype, jnp.integer):
    group_sizes = group_sizes.astype(jnp.int32)

Type guard

def is_int_array(a) -> bool:
    return dtypes.issubdtype(np.asarray(a).dtype, np.integer)

Prevention

When it happens

Trigger: 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.

Common situations: Computing token-per-expert counts with float ops (softmax, division) and forgetting to cast; loading counts from a float numpy array or CSV.

Related errors


AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27). Data as JSON: /api/errors/26deec772604e08a. Report an issue: GitHub.