jax-ml/jax · error · ValueError

Stride {s} is not divisible by {d} (tile size = {t}). Stride

Error message

Stride {s} is not divisible by {d} (tile size = {t}). Strides: {strides}, tiling: {tiling}

What it means

For nested tiles the stride of each dimension must be divisible by (previous stride * next tile size); otherwise the tiles overlap non-uniformly and the tiled layout cannot be built.

Source

Thrown at jax/experimental/mosaic/gpu/dialect_lowering.py:1060

  tiled_ordered_strides_and_tiling = sorted(
      tiled_strides_and_tiling, reverse=True)

  to_ordered = lambda i: tiled_ordered_strides_and_tiling.index(tiled_strides_and_tiling[i])
  from_ordered = lambda i: tiled_strides_and_tiling.index(tiled_ordered_strides_and_tiling[i])

  ordered_tiling = [tiling[from_ordered(i)] for i in range(len(tiling))]
  ordered_tiled_strides = [tiled_strides[from_ordered(i)] for i in range(len(tiling))]

  ordered_tiled_tiling_strides = [1]
  for t in reversed(ordered_tiling):
    ordered_tiled_tiling_strides.append(ordered_tiled_tiling_strides[-1] * t)

  prev_s = ordered_tiled_strides[-1]
  for s, t in zip(ordered_tiled_strides[:-1][::-1], ordered_tiling[1:][::-1], strict=True):
    d = prev_s * t
    prev_s = s
    if s % d != 0:
      raise ValueError(
          f"Stride {s} is not divisible by {d} (tile size = {t}). "
          f"Strides: {strides}, tiling: {tiling}"
      )
    ordered_tiled_tiling_strides.append(s // d * ordered_tiled_tiling_strides[-1])

  ordered_tiled_tiling_strides.reverse()

  return (
      *untiled_strides,
      *[ordered_tiled_tiling_strides[to_ordered(i)] for i in range(len(tiling))],
      *[ordered_tiled_tiling_strides[len(tiling) + to_ordered(i)] for i in range(len(tiling))]
  )


def transform_type(
    ref_ty: ir.MemRefType,
    transforms: tuple[lc.MemRefTransform, ...] | ir.ArrayAttr,
) -> ir.MemRefType:

View on GitHub (pinned to 1e1c6a8fc0)

Solutions

  1. Pad dimensions so strides are products of tile sizes (contiguous nested tiles)
  2. Choose tile sizes that divide the dimension strides
  3. Avoid interleaving untiled/padding elements inside tiled dimensions

Example fix

// before
tile_strides((2048, 48, 1), (8, 4))  # 48 not divisible by 4
// after
pad dim to 64 -> tile_strides((2048, 64, 1), (8, 4))
Defensive patterns

Strategy: validation

Validate before calling

# check nested tile stride divisibility before calling
# each inner stride must be divisible by (prev_stride * next_tile)
assert all(s % (ps * t) == 0 for s, ps, t in zip(inner_strides[:-1][::-1], inner_strides[1:][::-1], tiling[1:][::-1]))

Prevention

When it happens

Trigger: tile_strides where an inner tiled stride is not a multiple of the enclosing stride times tile size, e.g. strides (2048, 48, 1) with tiling (8, 4) since 48 % (1*4) != 0.

Common situations: Padding a tensor so inner dims are not multiples of the tile size (stride 48 with tile 4), producing misaligned nested tiles.

Related errors


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