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
- Pad dimensions so strides are products of tile sizes (contiguous nested tiles)
- Choose tile sizes that divide the dimension strides
- 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
- Pad tensor dims to multiples of tile sizes so strides stay products of tile sizes
- Avoid padding inside tiled dimensions
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
- Offset {i} is not divisible by tile size {t}
- Strides {strides} have lower rank than tiling {tiling}
- Can not tile strides when tiled dimensions have been transpo
- Transfer of {total_bits} bits is not divisible by {8 * utils
- `strides` must contain only 1s.
AI-assisted analysis of jax-ml/jax@1e1c6a8fc0 (2026-08-27).
Data as JSON: /api/errors/58c554958bbabbc0.
Report an issue: GitHub.