Skip to content

MetalLinear returns wrong outputs for every batch row after the first (3D inputs) #49337

Description

@francip

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

  • The official example scripts
  • My own modified scripts

Tasks

  • An officially supported task in the examples folder (such as GLUE/SQuAD, ...)
  • My own task or dataset (give details below)

Reproduction

Repro steps

  1. Load LFM2.5-230M with MetalConfig on MPS (M4 Pro, torch 2.14.0 native Metal kernel build)
  2. Perform a batched run.
  3. 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).

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

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions