A top-k choice is a control value
MLX-LM v0.32.0 includes PR #1930. The patch puts mx.stop_gradient on argpartition indices before scores are gathered, and its author adds narrow backward tests. GitHub release publication remains unknown; PyPI records a distinct October 1 package upload.
The argpartition reference specifies integer indices and undefined order within a partition. The stop_gradient reference specifies unchanged values with gradient flow blocked. These current MLX 0.32.3 docs describe operations, not a compatibility test for every installation.
In the original example below, score values are differentiable, while selected integer positions control which values are gathered. Detach those positions and gather from the original score tensor; gradients can still reach the selected values within a fixed selected region. A tie or rank crossing at the top-k boundary can change that region. Check finite backward computation and the analytic values for a fixed selection, without demanding a stable tie winner or smooth derivatives across a selection change.
- 1trainable gate scores
- 2argpartition top-k
- 3integer indices
- 4stop_gradient
- 1original gate scores + detached indices
- 2take_along_axis
- 3selected scores
- 4loss
- 5gradients to gate scores
- 1score tie or rank crossing
- 2selected index set may change
- 3fixture checks boundary, not a fixed winner
One deliberately small, unexecuted fixture
This original exercise remains unexecuted. PyPI metadata requires Python 3.11+ and MLX 0.32.2+ on Darwin for mlx-lm==0.32.0. Use that Apple-silicon path and record the resolved MLX version; the minimum is not a reproducible dependency pin. If the import or version check fails, stop and diagnose it.
from importlib.metadata import version
import mlx.core as mx
k = 2
def loss(gates):
inds = mx.argpartition(-gates, kth=k - 1, axis=-1)[..., :k]
inds = mx.stop_gradient(inds)
chosen = mx.take_along_axis(gates, inds, axis=-1)
return mx.sum(chosen * chosen)
assert version("mlx-lm") == "0.32.0"
print({"mlx_lm": version("mlx-lm"), "mlx": version("mlx")})
# A non-tied baseline proves more than shape/dtype/finite checks:
# detached all-zero values could satisfy those weaker checks.
baseline = mx.array([[1.01, 1.0, 0.2, -0.1]], dtype=mx.float32)
baseline_grad = mx.grad(loss)(baseline)
expected_baseline = mx.array([[2.02, 2.0, 0.0, 0.0]], dtype=mx.float32)
mx.eval(baseline_grad)
assert baseline_grad.shape == baseline.shape
assert baseline_grad.dtype == baseline.dtype
assert bool(mx.all(mx.isfinite(baseline_grad)))
assert bool(mx.allclose(baseline_grad, expected_baseline))
# The top-k boundary tie must not assume which equal middle score wins.
boundary = mx.array([[1.0, 0.2, 0.2, -0.1]], dtype=mx.float32)
boundary_grad = mx.grad(loss)(boundary)
mx.eval(boundary_grad)
assert bool(mx.all(mx.isfinite(boundary_grad)))
assert bool(mx.allclose(boundary_grad[:, [0, 3]],
mx.array([[2.0, 0.0]], dtype=mx.float32)))
assert bool(mx.allclose(mx.sum(boundary_grad), mx.array(2.4, dtype=mx.float32)))
assert bool(mx.allclose(mx.sort(boundary_grad[0, 1:3]),
mx.array([0.0, 0.4], dtype=mx.float32)))Both tensors have shape (1, 4) and use float32. The non-tied baseline selects the first two scores, so differentiating the squared selected scores analytically gives [[2.02, 2.0, 0.0, 0.0]]; this expectation is reasoning about the fixture, not a measurement. Shape, dtype, and finite checks alone would also pass for detached all-zero values, so they are insufficient.
The boundary case has [1.0, 0.2, 0.2, -0.1] with k=2. The first score must contribute 2.0, the last 0.0, and the gradient sum 2.4. Exactly one of the tied middle positions contributes 0.4; sorting only those two gradient values checks [0.0, 0.4] without asserting a winner or the order returned by argpartition. Record MLX/MLX-LM versions, macOS version, device, exact tensor values, k, and any exception. A missing package, version mismatch, unsupported device, non-finite gradient, or changed contract is a failed fixture. Stop there and inspect the pinned source; no fallback backend makes that result valid.
The upstream tests are author-maintained evidence, not an independent correctness or convergence result. This fixture does not test a full Mistral4 or MiniMax model, optimizer state, distributed training, quantization, or model quality. Its cost is local compute and memory only after an operator installs dependencies; it has no paid API call. Treat downloaded code and weights as separate trust and license decisions. MLX-LM's MIT license covers its repository code, not every model or dataset used with it.
A dated Apple-silicon context, without borrowing a performance claim
Apple's July 17, 2024 ICML update described an MLX demonstration of on-device inference and training on Apple silicon. That dated context does not benchmark MLX-LM 0.32 or establish this routing fix. The repair evidence comes from the 2026 patch; the small exercise remains a separate, unexecuted teaching example.
MENTAL MODEL / REASONING ORDER
From an announcement to your own decision.
Compare the announcement with the conditions in the paper and official documentation.
Sources
Publication dates belong to the source; access dates record when it was checked. Community observations are separate from official statements.