Skip to content

[BUG] mlx.sort return value incorrect on GPU #4538

Description

@kvcache670

☑️ I understand it is strictly prohibited to use AI to write issues.

Describe the bug
For a special shape (65537, 65536), and dtype uint8, the sort result is incorrect. For example, the row 0 is not sorted row 0 but other row. Check the reproduce.

To Reproduce

Requires an Apple Silicon GPU and approximately 18 GB of free memory. Save the following as repro.py.

python3 -m venv .venv
source .venv/bin/activate
python -m pip install "mlx==0.32.2" numpy
python repro.py
# mx.sort along the last axis of a row-contiguous uint8 array with a little more than 2^32 elements.
# Needs about 18 GB of free memory.
import mlx.core as mx
import numpy as np

R, C = 65537, 65536                                  # R * C = 2^32 + 65536
a = np.random.default_rng(3).integers(0, 251, size=(R, C), dtype=np.uint8)
x = mx.array(a)
y = mx.sort(x, axis=-1)
mx.eval(y)

print("mlx", mx.__version__, "|", mx.default_device(), "|", f"{R * C:,} elements")
for r in (0, 1, R // 2, R - 1):
    ok = np.array_equal(np.array(y[r]), np.sort(a[r]))
    print(f"row {r:5d}: {'ok' if ok else 'WRONG'}")
print("output row 0 == sort(input row R-1):", np.array_equal(np.array(y[0]), np.sort(a[R - 1])))
assert np.array_equal(np.array(y[0]), np.sort(a[0])), "sort(x)[0] is not the sorted row 0 of x"

Expected behavior
Each output row should equal np.sort(a[row]). Sorting along the last axis must not mix data between rows.

Desktop (please complete the following information):

  • OS Version: macOS 26.6.1 (25G76)
  • Hardware: Apple M1 Ultra
  • MLX: 0.32.2 (pip wheel)
  • Unified memory: 128 GB

Additional context

Here is the analysis of AI (I know there's AI policy, I want to post it here to better explain the analysis).

The row offset at row 65536 is 65536 * 65536 = 2^32, which suggests a 32-bit offset overflow. In sort.h at the inspected revision, the expressions tid.y * in_stride_segment_axis and tid.y * out_stride_segment_axis use a uint3 threadgroup index and int strides. This is a suspected cause; the reproduction does not establish which sort or merge kernel performs the incorrect write.

The recorded run uses the 0.32.2 wheel. The same expressions were found in source at 59d600b, but that revision was not built or run for this report. CPU sorting at this size has not been tested.

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