Skip to content

[BUG] mx.split view data_size under-reported → elementwise int32 index overflow → silent corruption of rows >= ceil(2^31/stride) #4548

Description

@fenglingai123

[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.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions