fix(kernel): avoid row-wise _scaled_mm stall on sm_89 with torch<2.12 (#243)

PyTorch < 2.12 runs row-wise FP8 _scaled_mm on sm_89 through a CUTLASS stream-K kernel whose launch ignored the current stream (pytorch/pytorch#177651, fixed by pytorch/pytorch@252bb4a in 2.12). FreeToken issues the fused per-tensor-FP8 projections (q/k/v and GDN qkv|z of the NVFP4 checkpoints) from a side stream, so on Ada every prefill of >= 256 tokens stalled the GPU and the worker hung or died (#182, #72, #220). Windows torch builds ship no row-wise kernel at all (#227).

Tensor-wise scaling is unaffected. Where row-wise is unsafe (sm_89 on torch < 2.12, or a probe on the default stream raises), a fused projection now runs one tensor-wise GEMM per part over its row slice and concatenates: the same W8A8 scheme (rel ~7e-4 to row-wise, accumulation order), one extra launch per part. The parts' row ranges come from the load-time weight_scale run-lengths; the decision and its probe run at load, never under CUDA-graph capture. FREETOKEN_FP8_ROWWISE_MM=0/1 forces either path for A/B.

Tested on RTX 4070 SUPER (sm_89), driver 591.86, torch 2.11.0+cu130, WSL2. Sweep over M on a side stream: row-wise stalls at M >= 256, the new path completes at every M. tests/kernels/test_fp8_pertensor_linear.py: the side-stream test fails on main (rc=124, 0/128 GEMMs complete) and passes here; the per-part path is compared directly against row-wise at M=1/4/64/300. Three pre-existing test_w8a8_matches_w8a8_reference cases miss the 1e-2 tolerance on this GPU (rel 0.0103-0.0107) on main and on this branch alike.

Assisted-by: Claude Fable 5
This commit is contained in:
Berni McCoy
2026-08-29 00:20:16 -04:00
committed by GitHub
parent a05c26543f
commit 58f4b9ec0e
2 changed files with 167 additions and 6 deletions
@@ -19,7 +19,9 @@ accurate than the previous ``weight.to(bf16) * scale`` materialization, which it
from __future__ import annotations
import functools
import os
import re
import torch
import triton
@@ -42,6 +44,64 @@ _TL_DTYPE = {torch.bfloat16: tl.bfloat16, torch.float16: tl.float16, torch.float
_USE_REF = os.environ.get("FREETOKEN_DEBUG_FP8_REF") == "1"
# Row-wise _scaled_mm on sm_89 with torch < 2.12 launches its CUTLASS stream-K kernel off the
# current stream (pytorch/pytorch#177651, fixed by pytorch/pytorch@252bb4a; #182/#72/#220), and
# some builds (Windows) have no row-wise kernel (#227). Fallback: one tensor-wise GEMM per part.
def _torch_version() -> tuple[int, int]:
m = re.match(r"(\d+)\.(\d+)", torch.__version__)
return (int(m.group(1)), int(m.group(2))) if m else (0, 0)
@functools.cache
def rowwise_scaled_mm_ok() -> bool:
"""Whether row-wise ``torch._scaled_mm`` may be issued from a side stream on this GPU.
Decided once per process, at load (never under graph capture). ``FREETOKEN_FP8_ROWWISE_MM=0/1``
forces the answer."""
forced = os.environ.get("FREETOKEN_FP8_ROWWISE_MM")
if forced in ("0", "1"):
return forced == "1"
if not torch.cuda.is_available():
return True
from freetoken.gpu_select import assigned_visible_gpu
idx = assigned_visible_gpu()
dev = torch.device("cuda", torch.cuda.current_device() if idx is None else idx)
if torch.cuda.get_device_capability(dev) == (8, 9) and _torch_version() < (2, 12):
return False
# Probe on the default stream (safe even where the launch ignores the current stream); a
# build without the row-wise kernel raises here instead of at the first forward.
try:
with torch.cuda.device(dev), torch.cuda.stream(torch.cuda.default_stream(dev)):
a = torch.zeros(16, 32, dtype=FP8, device=dev)
b = torch.zeros(32, 32, dtype=FP8, device=dev)
torch._scaled_mm(
a, b.t(), scale_a=torch.ones(16, 1, device=dev),
scale_b=torch.ones(1, 32, device=dev), out_dtype=torch.bfloat16,
)
torch.cuda.synchronize(dev)
except RuntimeError:
return False
return True
def weight_scale_segments(weight_scale: torch.Tensor) -> list[tuple[int, int]]:
"""``[start, end)`` row ranges over which ``weight_scale`` is constant (the fused parts).
Syncs; call at load."""
s = weight_scale.detach().reshape(-1).float().cpu()
change = (torch.nonzero(s[1:] != s[:-1]).flatten() + 1).tolist()
bounds = [0, *change, s.numel()]
return list(zip(bounds[:-1], bounds[1:]))
_MAX_SEGMENTS = 8 # q/k/v = 3, GDN qkv|z = 2; a genuine per-row scale stays W8A16 instead
def _segments_w8a8_ok(segments: list[tuple[int, int]]) -> bool:
"""cuBLASLt needs 16-row aligned fp8 operands; more parts than a fused projection has
means a genuine per-row scale."""
return 0 < len(segments) <= _MAX_SEGMENTS and all((e - s) % 16 == 0 for s, e in segments)
# ======================================================================================
# Decode (M==1) split-K GEMV: raw fp8 x bf16 reduction in fp32, per-row scale at reduce.
# ======================================================================================
@@ -222,6 +282,7 @@ def _static_quant(a: torch.Tensor, input_scale: torch.Tensor) -> torch.Tensor:
def _scaled_mm(
a: torch.Tensor, weight: torch.Tensor, weight_scale: torch.Tensor,
input_scale: torch.Tensor, uniform_scale: bool, out_dtype: torch.dtype,
scale_segments: list[tuple[int, int]] | None = None,
) -> torch.Tensor:
"""``a @ (weight_fp8 * weight_scale)^T`` as a W8A8 cuBLASLt GEMM.
@@ -233,14 +294,26 @@ def _scaled_mm(
tensor-wise path. A fused projection, whose ``weight_scale`` is piecewise-constant
because each part carries its own scalar, takes the row-wise path -- that keeps every
part's scale exact, where vLLM/SGLang instead requantize the parts onto a shared maximum
and eat the precision loss. Row-wise costs ~4% here (5.56 ms vs 5.39 ms per step)."""
and eat the precision loss. Row-wise costs ~4% here (5.56 ms vs 5.39 ms per step).
Where row-wise is unsafe (:func:`rowwise_scaled_mm_ok`) ``scale_segments`` is passed and
each part runs its own tensor-wise GEMM over ``weight[s:e]`` (still stride-only), outputs
concatenated: the same W8A8 scheme, not bit-identical (accumulation order differs)."""
qa = _static_quant(a, input_scale)
wt = weight.t() # [N, K] row-major -> [K, N] column-major, stride-only
sa = input_scale.reshape(())
if uniform_scale:
return torch._scaled_mm(
qa, wt, scale_a=input_scale.reshape(()), scale_b=weight_scale[0].reshape(()),
out_dtype=out_dtype,
qa, wt, scale_a=sa, scale_b=weight_scale[0].reshape(()), out_dtype=out_dtype,
)
if scale_segments is not None:
return torch.cat([
torch._scaled_mm(
qa, weight[s:e].t(), scale_a=sa, scale_b=weight_scale[s].reshape(()),
out_dtype=out_dtype,
)
for s, e in scale_segments
], dim=1)
return torch._scaled_mm(
qa, wt,
scale_a=input_scale.reshape(1, 1).expand(a.shape[0], 1).contiguous(),
@@ -254,9 +327,11 @@ def fp8_pertensor_linear(
bias: torch.Tensor | None = None,
input_scale: torch.Tensor | None = None,
uniform_scale: bool = False,
scale_segments: list[tuple[int, int]] | None = None,
) -> torch.Tensor:
"""``y = x @ (weight_fp8 * weight_scale)^T``. ``weight`` [N, K] fp8-e4m3, ``weight_scale``
[N] fp32 (per output row).
[N] fp32 (per output row). ``scale_segments``: the fused parts' row ranges, precomputed at
load by the layer; derived here (with a sync) when omitted and needed.
Whether the activation is quantized is a property of the *deployment*, never of the batch:
with ``input_scale`` on sm_89+ every M runs W8A8, otherwise every M runs W8A16 (split-K
@@ -266,12 +341,19 @@ def fp8_pertensor_linear(
SGLang likewise run one scheme across all M on any GPU with FP8 tensor cores."""
*lead, K = x.shape
N = weight.shape[0]
w8a8 = input_scale is not None and e4m3_native()
segments = None
if w8a8 and not uniform_scale and not rowwise_scaled_mm_ok():
segments = scale_segments if scale_segments is not None else weight_scale_segments(weight_scale)
if not _segments_w8a8_ok(segments):
w8a8 = False # W8A16 below is exact for any per-row scale and never calls _scaled_mm
if _USE_REF: # numeric-reference fallback (debug / A-B)
w = weight.to(x.dtype) * weight_scale.to(x.dtype)[:, None]
out = (x.reshape(-1, K) @ w.t()).reshape(*lead, N)
elif input_scale is not None and e4m3_native():
elif w8a8:
out = _scaled_mm(
x.reshape(-1, K), weight, weight_scale, input_scale, uniform_scale, x.dtype,
scale_segments=segments,
).reshape(*lead, N)
elif x.numel() // K == 1:
out = _gemv(x.reshape(K), e4m3_kernel_view(weight), weight_scale, x.dtype).reshape(*lead, N)
@@ -306,6 +388,7 @@ class Fp8PerTensorLinear(BaseOP):
# reflective state_dict/load_state_dict skip it entirely on checkpoints without one.
self.input_scale: torch.Tensor | None = None
self._uniform_scale = False
self._scale_segments: list[tuple[int, int]] | None = None
def load_state_dict(self, state_dict, *, prefix: str = "", _internal: bool = False) -> None:
# Taken out before BaseOP's reflective pass (so it is not an "unexpected key") and
@@ -318,11 +401,15 @@ class Fp8PerTensorLinear(BaseOP):
# only piecewise-constant, so decide once here rather than syncing on every forward.
scale = self.weight_scale
self._uniform_scale = bool((scale == scale[0]).all().item())
# Segments for the per-part path; decide row-wise safety now, not under graph capture.
self._scale_segments = None if self._uniform_scale else weight_scale_segments(scale)
if self.input_scale is not None and not self._uniform_scale:
rowwise_scaled_mm_ok()
def forward(self, x: torch.Tensor) -> torch.Tensor:
return fp8_pertensor_linear(
x, self.weight, self.weight_scale, self.bias,
self.input_scale, self._uniform_scale,
self.input_scale, self._uniform_scale, scale_segments=self._scale_segments,
)
@@ -341,4 +428,6 @@ __all__ = [
"Fp8PerTensorLinear",
"Fp8PerTensorColMerged",
"fp8_pertensor_linear",
"rowwise_scaled_mm_ok",
"weight_scale_segments",
]
@@ -9,6 +9,10 @@ to fp8 and is held to a reference that applies the same quantization.
from __future__ import annotations
import subprocess
import sys
import textwrap
import pytest
import torch
@@ -125,3 +129,71 @@ def test_layer_load_marks_uniform_scale_and_optional_input_scale():
# a reload must not trip over the input_scale it kept from the first load
single.load_state_dict({"weight": w8, "weight_scale": flat})
assert single.input_scale is None
@pytest.mark.skipif(not e4m3_native(), reason="torch._scaled_mm needs sm_89+")
@pytest.mark.parametrize("M", [1, 4, 64, 300])
def test_per_part_path_matches_rowwise(M: int, monkeypatch):
"""Where row-wise ``_scaled_mm`` is unsafe a fused projection runs one tensor-wise GEMM per
part instead. Same scheme, so the two paths agree up to accumulation order (~7e-4)."""
import freetoken.kernel.triton.fp8_pertensor_linear as mod
K, part_rows = 2048, [1024, 256, 256]
w8, scale = _quant_parts(part_rows, K, seed=M)
x = torch.randn(M, K, device=DEV, dtype=torch.bfloat16)
input_scale = (x.abs().max().float() / 448.0).reshape(())
monkeypatch.setattr(mod, "rowwise_scaled_mm_ok", lambda: True)
y_row = mod.fp8_pertensor_linear(x, w8, scale, None, input_scale, False)
monkeypatch.setattr(mod, "rowwise_scaled_mm_ok", lambda: False)
y_part = mod.fp8_pertensor_linear(x, w8, scale, None, input_scale, False)
rel = ((y_part.float() - y_row.float()).norm() / y_row.float().norm()).item()
assert rel < 2e-3, rel
@pytest.mark.skipif(not e4m3_native(), reason="torch._scaled_mm needs sm_89+")
def test_fused_layer_forward_on_a_side_stream_completes():
"""Regression for #182 / #72 / #220: on sm_89 with torch < 2.12 a fused FP8 projection's
row-wise ``_scaled_mm`` issued from a non-default stream stalls the GPU (PyTorch's
CUTLASS row-wise kernel ignored the current stream; fixed upstream in 2.12). The layer
must take a path that completes on every supported build. Runs in a subprocess so a
stall fails the test instead of hanging the session."""
script = textwrap.dedent("""
import os, time, torch
from freetoken.kernel.triton.fp8_pertensor_linear import FP8, Fp8PerTensorColMerged
torch.manual_seed(0)
K, parts = 2048, [8192, 512, 512] # a prefill-sized fused qkv
w8 = (torch.randn(sum(parts), K, device="cuda") * 8).clamp(-448, 448).to(FP8)
scale = torch.cat([torch.full((p,), 0.01 * (i + 1), device="cuda")
for i, p in enumerate(parts)])
layer = Fp8PerTensorColMerged(K, parts)
layer.load_state_dict({"weight": w8, "weight_scale": scale,
"input_scale": torch.tensor(0.02, device="cuda")})
x = torch.randn(2010, K, device="cuda", dtype=torch.bfloat16) # #182 shape
torch.cuda.synchronize()
stream = torch.cuda.Stream()
events = []
with torch.cuda.stream(stream):
for _ in range(128):
layer.forward(x)
ev = torch.cuda.Event()
ev.record(stream)
events.append(ev)
deadline = time.monotonic() + 30
done = 0
while time.monotonic() < deadline:
while done < len(events) and events[done].query():
done += 1
if done == len(events):
print("completed", done, flush=True)
os._exit(0)
time.sleep(0.05)
print("stalled at", done, "of", len(events), flush=True)
os._exit(124) # a normal exit would wait on the stuck kernel
""")
proc = subprocess.run(
[sys.executable, "-c", script], capture_output=True, text=True, timeout=300,
)
assert proc.returncode == 0, f"rc={proc.returncode}\n{proc.stdout}\n{proc.stderr[-2000:]}"