System Info
Apple M4 Pro, 64 GiB, macOS 26
transformers main @ 8073e6d
torch 2.14.0 (native MPS), kernels 0.17.2
Who can help?
@SunMarc (quantization), @vasqu / @drbh (kernels)
Information
Tasks
Reproduction
Repro steps
- Load LFM2.5-230M with MetalConfig on MPS (M4 Pro, torch 2.14.0 native Metal kernel build)
- Perform a batched run.
- Observe that every row after the first comes out as garbage, and some rows have non-finite logits.
Standalone repro:
uv venv .venv --python 3.12
uv pip install --python .venv/bin/python torch==2.14.0 transformers==5.16.1 kernels==0.16.1
uv pip install --python .venv/bin/python torch==2.14.0 kernels "transformers @ git+https://github.com/huggingface/transformers@main"
.venv/bin/python repro.py
repro.py:
import torch
from transformers.integrations.metal_quantization import (
MetalLinear, _affine_quantize_tensor,
)
assert torch.backends.mps.is_available()
torch.manual_seed(42)
K, N, GROUP, BITS = 512, 256, 64, 8
w, s, b = _affine_quantize_tensor(torch.randn(N, K), GROUP, BITS)
layer = MetalLinear(K, N, bits=BITS, group_size=GROUP).to("mps")
with torch.no_grad():
layer.weight.copy_(w)
layer.scales.copy_(s)
layer.qbiases.copy_(b)
x = torch.randn(4, 3, K, device="mps", dtype=torch.bfloat16)
with torch.inference_mode():
actual = layer(x)
flat = layer(x.reshape(-1, K)).reshape(4, 3, N)
rows = torch.stack([layer(row) for row in x])
torch.mps.synchronize()
print("3D finite by batch:", torch.isfinite(actual).all(-1).all(-1).tolist())
print("flattened finite:", bool(torch.isfinite(flat).all()))
print("3D vs individual max abs error by batch:",
(actual.float() - rows.float()).abs().amax((1, 2)).tolist())
print("flattened vs individual max abs error:",
(flat.float() - rows.float()).abs().max().item())
torch.testing.assert_close(flat, rows, rtol=0.02, atol=0.125)
# Fails without the integration-side flatten/restore fix.
torch.testing.assert_close(actual, rows, rtol=0.02, atol=0.125)
Details:
MetalLinear.forward passes the raw N-D input to affine_qmm_t. When the input is 3D (batch, seq, hidden), the Metal kernel only computes batch element 0 correctly. Every later row is garbage, so batched generate() with MetalConfig silently corrupts every row except the first.
Batch size 1 is fine, unquantized BF16 are fine at every batch size, MLX / MLX-LM 0.32.2 / 0.31.3 comparison of the same model is fine.
This is the same bug as #47331, which went stale and was auto-closed without a fix. The kernel-side fix (kernels-community#1038) was closed by a bot and never merged. The bug is still present in latest main.
Expected behavior
Each row of a batched MetalLinear forward should equal that row run alone. At the model level, batched greedy generate() with MetalConfig should give each row the same text as batch size 1, as unquantized bf16 already does.
Proposed fix
MetalLinear.forward should flatten the input to 2D before affine_qmm_t and reshape the output afterwards. This is a two-line change and it does not depend on a kernel release. I can open a PR with the fix and regression tests (a mocked CPU test for shape handling, plus an MPS test that compares batched and per-row outputs for 4/8 bits, fp16/bf16 and with/without bias).
System Info
Apple M4 Pro, 64 GiB, macOS 26
transformers main @ 8073e6d
torch 2.14.0 (native MPS), kernels 0.17.2
Who can help?
@SunMarc (quantization), @vasqu / @drbh (kernels)
Information
Tasks
examplesfolder (such as GLUE/SQuAD, ...)Reproduction
Repro steps
Standalone repro:
repro.py:
Details:
MetalLinear.forwardpasses the raw N-D input toaffine_qmm_t. When the input is 3D (batch, seq, hidden), the Metal kernel only computes batch element 0 correctly. Every later row is garbage, so batchedgenerate()withMetalConfigsilently corrupts every row except the first.Batch size 1 is fine, unquantized BF16 are fine at every batch size, MLX / MLX-LM 0.32.2 / 0.31.3 comparison of the same model is fine.
This is the same bug as #47331, which went stale and was auto-closed without a fix. The kernel-side fix (kernels-community#1038) was closed by a bot and never merged. The bug is still present in latest main.
Expected behavior
Each row of a batched
MetalLinearforward should equal that row run alone. At the model level, batched greedygenerate()withMetalConfigshould give each row the same text as batch size 1, as unquantized bf16 already does.Proposed fix
MetalLinear.forwardshould flatten the input to 2D beforeaffine_qmm_tand reshape the output afterwards. This is a two-line change and it does not depend on a kernel release. I can open a PR with the fix and regression tests (a mocked CPU test for shape handling, plus an MPS test that compares batched and per-row outputs for 4/8 bits, fp16/bf16 and with/without bias).