[BUG] mx.split view data_size under-reported → elementwise int32 index overflow → silent corruption of rows >= ceil(2^31/stride)
MLX version: 0.32.2 (also checked current main: the code path is unchanged)
Platform: macOS 26.x arm64, Apple M4 Max, Metal backend
Related: #3866 (closed, asked for an mlx-only minimal repro — this is one), #3784 (open, same size-gated silent-corruption family)
Summary
When the source tensor has more than 2^31 elements, mx.split returns views whose data_size() is the logical element count (shape[0] * shape[1]), not the real buffer span ((shape[0]-1) * src_stride[0] + ...). Downstream elementwise kernels (binary/unary) select int32 vs int64 element indexing via data_size() > INT32_MAX, so a view with logical size 1.12e9 but actual source offsets up to 2.24e9 takes the int32 path and silently reads wrong addresses for rows >= ceil(2^31 / src_stride).
No error, no crash — the output simply contains zeros/garbage from that row on. This is data corruption that survives mx.eval, .copy(), astype, and reshaping of the source.
Minimal repro (pure mlx python, deterministic, 100%)
repro_issue.py (attached / below) — ~50 lines, single lazy chain, one eval:
y = mx.quantized_matmul(x, wq, sc, bs, transpose=True, group_size=64, bits=8, mode="affine") # [78166, 28672] = 2.24e9 elems
g, u = mx.split(y, 2, axis=1) # views: shape [78166, 14336], strides [28672, 1]
act = ((g * mx.sigmoid(g)) * u).astype(mx.bfloat16)
mx.eval(act)
# compare against independent CPU float64 reference on a window straddling row 74899
Result (unfixed): every row >= 74898 in the output is zero/corrupt; rows below are bit-exact. First bad row = ceil(2^31 / 28672) - 1 = 74898. Deterministic across runs.
Result (with the one-line fix below): all rows match the CPU reference.
Root cause
Split::eval in mlx/backend/common/common.cpp:
auto compute_new_flags = [](const auto& shape, const auto& strides,
size_t in_data_size, auto flags) {
size_t data_size = 1;
...
if (strides[i] > 0) {
data_size *= shape[i]; // ← LOGICAL size, not buffer span
}
...
};
array.h documents data_size as "the size (in elements) of the underlying buffer the array points to … last_offset − first_offset". For a contiguous array the logical product equals the span, so this is usually harmless. For a split view (shape [M, N/2], source strides [N, 1]) they diverge:
| quantity |
value |
vs 2^31 |
logical M*N/2 |
1.12e9 |
below → picks int32 kernels |
true span (M-1)*N + 1 |
2.24e9 |
above → int32 offsets overflow |
Instrumented proof (C++ probe, post-eval):
unfixed: y.g shape=[78166,14336] strides=[28672,1] data_size=1120587776 ← wrong
fixed: y.g shape=[78166,14336] strides=[28672,1] data_size=2241161216 ← correct span
Consumers that then mis-dispatch: binary.cpp:95, unary.cpp:40, ternary.cpp:43 — all gate on in.data_size() > INT32_MAX. Kernels instantiated with IdxT=int (binary.metal:29 g2_*) compute elem_to_loc_2<int>(elem, strides) → elem.y * (int)28672 overflows at row 74899.
Why astype-only consumption appears clean: the General copy path computes offsets with int64 accumulation independently of data_size in some branches — but binary/unary on the view always corrupt. Our production failure hit exactly this: SwiGLU (silu(g) * u) after a split of a 2.24e9-element activation.
Proposed fix (verified)
In Split::eval, hand the true buffer span to copy_shared_buffer (leave compute_new_flags's contiguity logic on the logical value — it compares against in.data_size()):
size_t buffer_span = 1;
const auto& out_shape = outputs[i].shape();
const auto& in_strides = in.strides();
for (size_t d = 0; d < out_shape.size(); d++) {
if (in_strides[d] > 0 && out_shape[d] > 1) {
buffer_span += (out_shape[d] - 1) * static_cast<size_t>(in_strides[d]);
}
}
const size_t true_span = data_size > buffer_span ? data_size : buffer_span;
outputs[i].copy_shared_buffer(in, in.strides(), new_flags, true_span, offset);
For contiguous inputs span == logical size → zero behavior change. For non-contiguous views, data_size() > INT32_MAX becomes true and every consumer correctly selects the int64 kernels.
Verified on M4 Max: repro turns clean; a battery of 5 consumer variants (slice/split × elementwise/copy × contiguous) all pass.
Honest observation we could not explain
With a random-normal materialized source of the same shape/dtype (mx.random.normal([M,N]).astype(bfloat16), eval'ed), the same split+SwiGLU chain does not corrupt; with a quantized_matmul output source it corrupts 100% deterministically — even after mx.eval(y) (materialized) and even after a deep copy (mx.array(y)), where the zeros turn into non-zero garbage (offsets still wrong, landing inside the buffer). Post-eval internal dumps of both sources look identical (shape/strides/flags/data_size). We report this as an observation, not a claim: the fix above is proven either way, since it corrects data_size for every split view. Happy to run further probes if maintainers want this pinned down.
Impact
Any model whose intermediate activations exceed 2^31 elements and that splits them along the feature axis (large-batch / long-context SwiGLU FFNs — exactly our production video-diffusion transformer at 243-frame 768p / 209-frame 1088p). Silent, deterministic, size-gated: below the threshold everything is perfect, above it the tail rows are garbage with no diagnostic. We spent 16 debugging runs across two weeks before pinning it down.
[BUG] mx.split view data_size under-reported → elementwise int32 index overflow → silent corruption of rows >= ceil(2^31/stride)
MLX version: 0.32.2 (also checked current main: the code path is unchanged)
Platform: macOS 26.x arm64, Apple M4 Max, Metal backend
Related: #3866 (closed, asked for an mlx-only minimal repro — this is one), #3784 (open, same size-gated silent-corruption family)
Summary
When the source tensor has more than 2^31 elements,
mx.splitreturns views whosedata_size()is the logical element count (shape[0] * shape[1]), not the real buffer span ((shape[0]-1) * src_stride[0] + ...). Downstream elementwise kernels (binary/unary) select int32 vs int64 element indexing viadata_size() > INT32_MAX, so a view with logical size 1.12e9 but actual source offsets up to 2.24e9 takes the int32 path and silently reads wrong addresses for rows>= ceil(2^31 / src_stride).No error, no crash — the output simply contains zeros/garbage from that row on. This is data corruption that survives
mx.eval,.copy(),astype, and reshaping of the source.Minimal repro (pure mlx python, deterministic, 100%)
repro_issue.py(attached / below) — ~50 lines, single lazy chain, one eval:Result (unfixed): every row
>= 74898in the output is zero/corrupt; rows below are bit-exact. First bad row =ceil(2^31 / 28672) - 1 = 74898. Deterministic across runs.Result (with the one-line fix below): all rows match the CPU reference.
Root cause
Split::evalinmlx/backend/common/common.cpp:array.hdocumentsdata_sizeas "the size (in elements) of the underlying buffer the array points to … last_offset − first_offset". For a contiguous array the logical product equals the span, so this is usually harmless. For a split view (shape[M, N/2], source strides[N, 1]) they diverge:M*N/2(M-1)*N + 1Instrumented proof (C++ probe, post-eval):
Consumers that then mis-dispatch:
binary.cpp:95,unary.cpp:40,ternary.cpp:43— all gate onin.data_size() > INT32_MAX. Kernels instantiated withIdxT=int(binary.metal:29 g2_*) computeelem_to_loc_2<int>(elem, strides)→elem.y * (int)28672overflows at row 74899.Why
astype-only consumption appears clean: the General copy path computes offsets with int64 accumulation independently ofdata_sizein some branches — butbinary/unaryon the view always corrupt. Our production failure hit exactly this: SwiGLU (silu(g) * u) after a split of a 2.24e9-element activation.Proposed fix (verified)
In
Split::eval, hand the true buffer span tocopy_shared_buffer(leavecompute_new_flags's contiguity logic on the logical value — it compares againstin.data_size()):For contiguous inputs span == logical size → zero behavior change. For non-contiguous views,
data_size() > INT32_MAXbecomes true and every consumer correctly selects the int64 kernels.Verified on M4 Max: repro turns clean; a battery of 5 consumer variants (slice/split × elementwise/copy × contiguous) all pass.
Honest observation we could not explain
With a random-normal materialized source of the same shape/dtype (
mx.random.normal([M,N]).astype(bfloat16), eval'ed), the same split+SwiGLU chain does not corrupt; with aquantized_matmuloutput source it corrupts 100% deterministically — even aftermx.eval(y)(materialized) and even after a deep copy (mx.array(y)), where the zeros turn into non-zero garbage (offsets still wrong, landing inside the buffer). Post-eval internal dumps of both sources look identical (shape/strides/flags/data_size). We report this as an observation, not a claim: the fix above is proven either way, since it correctsdata_sizefor every split view. Happy to run further probes if maintainers want this pinned down.Impact
Any model whose intermediate activations exceed 2^31 elements and that splits them along the feature axis (large-batch / long-context SwiGLU FFNs — exactly our production video-diffusion transformer at 243-frame 768p / 209-frame 1088p). Silent, deterministic, size-gated: below the threshold everything is perfect, above it the tail rows are garbage with no diagnostic. We spent 16 debugging runs across two weeks before pinning it down.