Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
Show all changes
27 commits
Select commit Hold shift + click to select a range
11e9007
add support for row-wise quanted input for grouped gemm
YangFei1990 Jul 23, 2026
d234cd8
doc change
YangFei1990 Jul 25, 2026
376b94c
implement for backward
YangFei1990 Jul 26, 2026
782835c
merge from main
YangFei1990 Jul 26, 2026
855f480
allow scaled_bias + frozen weights in prequant path
YangFei1990 Jul 26, 2026
8ef7e7e
allow bias + frozen weight path in prequantized input
YangFei1990 Jul 26, 2026
6015c91
use .copy to create group tensor
YangFei1990 Jul 30, 2026
49ba134
merge from main
YangFei1990 Jul 30, 2026
12f1dbf
move the implementation to c++ layer
YangFei1990 Jul 30, 2026
14fedf0
Merge branch 'main' into mxfp8_input_groupedgemm
YangFei1990 Jul 30, 2026
e3990dd
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jul 30, 2026
c1e4663
bug fix for operation ordering
YangFei1990 Jul 31, 2026
a349d0d
add comprehensive tests
YangFei1990 Jul 31, 2026
a3219de
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jul 31, 2026
7c850ef
Merge branch 'main' into mxfp8_input_groupedgemm
YangFei1990 Jul 31, 2026
d8e8b1c
Merge branch 'main' into mxfp8_input_groupedgemm
YangFei1990 Aug 4, 2026
3bbf778
rename to group_requantize
YangFei1990 Aug 5, 2026
a677627
merge from main
YangFei1990 Aug 6, 2026
e890104
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 6, 2026
76662c2
rename func with inplace tag
YangFei1990 Aug 6, 2026
f2ff2fc
resolve merge
YangFei1990 Aug 6, 2026
a056c6d
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 6, 2026
ce2e5a3
refactor the requantize function
YangFei1990 Aug 6, 2026
8a1c89f
merge from main
YangFei1990 Aug 6, 2026
226405d
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Aug 6, 2026
3032804
minor fix for test
YangFei1990 Aug 7, 2026
791abaa
Merge branch 'main' into mxfp8_input_groupedgemm
YangFei1990 Aug 7, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
363 changes: 363 additions & 0 deletions tests/pytorch/test_grouped_mlp.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
import torch

import transformer_engine.pytorch as te
from transformer_engine.pytorch.constants import TE_DType
from transformer_engine.pytorch.ops.fused.grouped_mlp import (
_cudnn_frontend_supports_grouped_gemm_srelu,
_cudnn_frontend_version_supported,
Expand Down Expand Up @@ -229,6 +230,23 @@ def make_reference_and_test_tensors(
return ref, test


class _InjectGrad(torch.autograd.Function):
"""Replace the gradient flowing into ``x`` with ``grad``.

Mirrors how an FP8 token dispatch delivers a pre-quantized ``GroupedTensor``
grad output: the downstream op simply returns one as its grad input.
"""

@staticmethod
def forward(ctx, x, grad): # pylint: disable=arguments-differ
ctx.injected_grad = grad
return x

@staticmethod
def backward(ctx, grad_output): # pylint: disable=arguments-differ
return ctx.injected_grad, None


class TestGroupedLinearOp:
"""Tests for advanced features with grouped linear basic op"""

Expand Down Expand Up @@ -436,6 +454,164 @@ def test_grouped_linear(
else:
assert b_test.grad is None

@staticmethod
def _make_rowwise_mxfp8_wire_input(
x_hp: torch.Tensor,
group_size: int,
split_sizes: torch.Tensor,
) -> "GroupedTensor":
"""Rowwise-only MXFP8 GroupedTensor with compact scales (FP8 dispatch wire format)."""
wire_quantizer = MXFP8Quantizer(
fp8_dtype=tex.DType.kFloat8E4M3, rowwise=True, columnwise=False
Comment thread
YangFei1990 marked this conversation as resolved.
Outdated
)
x_wire = tex.group_quantize(
x_hp, wire_quantizer, group_size, split_sizes.to(dtype=torch.int64)
)
# Sanity: the wire tensor is rowwise-only with unswizzled scales.
assert x_wire.columnwise_data is None
assert not x_wire._with_gemm_swizzled_scales
return x_wire

@pytest.mark.parametrize("weight_requires_grad", (False, True))
def test_grouped_linear_prequantized_mxfp8_input(
self,
*,
group_size: int = 4,
weight_shape: tuple[int, int] = (256, 256),
split_alignment: int = 128,
dtype: torch.dtype = torch.bfloat16,
device: torch.device = "cuda",
weight_requires_grad: bool,
) -> None:
"""Rowwise-only MXFP8 GroupedTensor input (FP8 token dispatch wire format).

The input arrives already rowwise-quantized with compact scales. The op
must feed the rowwise data to the forward GEMM as-is and manufacture the
columnwise copy for the wgrad GEMM. The reference run consumes the
*dequantized* wire tensor (the only data a layer can see after FP8
dispatch) on the normal quantize-from-BF16 path; because MXFP8 requant
is idempotent along the rowwise axis and both paths derive the
columnwise copy from the same dequantized data, the two runs must match
bit-for-bit.
"""
if not mxfp8_available:
pytest.skip(reason_for_no_mxfp8)
maybe_skip_quantization("mxfp8", dims=weight_shape, device=device, dtype=dtype)

# Split sizes (including an empty group)
split_sizes = [split_alignment * i for i in range(group_size)]
random.shuffle(split_sizes)
split_sizes = torch.tensor(split_sizes, dtype=torch.int, device=device)

out_features, in_features = weight_shape
total_tokens = int(split_sizes.sum().item())
in_shape = (total_tokens, in_features)

# Wire-format input and its exact dequantization (the reference input).
x_hp = torch.rand(in_shape, dtype=dtype, device=device) - 0.5
x_wire = self._make_rowwise_mxfp8_wire_input(x_hp, group_size, split_sizes)
x_ref = tex.group_dequantize(x_wire, TE_DType[dtype]).rowwise_data.view(in_shape)

dy = torch.rand((total_tokens, out_features), dtype=dtype, device=device) - 0.5
recipe = make_recipe("mxfp8")
op = te.ops.GroupedLinear(
group_size, in_features, out_features, bias=False, device=device, dtype=dtype
)
with torch.no_grad():
for param in op.parameters():
param.requires_grad_(requires_grad=weight_requires_grad)

def _run(x):
with te.autocast(enabled=True, recipe=recipe):
y = op(x, split_sizes)
wgrads = []
if weight_requires_grad:
y.backward(dy)
for group_idx in range(group_size):
weight = getattr(op, f"weight{group_idx}")
wgrads.append(weight.grad.detach().clone())
weight.grad = None
return y.detach(), wgrads

y_ref, wgrads_ref = _run(x_ref)
y_test, wgrads_test = _run(x_wire)

# Bit-exact match expected (identical quantized inputs and kernels).
torch.testing.assert_close(y_test, y_ref, rtol=0, atol=0)
for wgrad_test, wgrad_ref in zip(wgrads_test, wgrads_ref):
torch.testing.assert_close(wgrad_test, wgrad_ref, rtol=0, atol=0)

@pytest.mark.parametrize("bias", (False, True))
def test_grouped_linear_prequantized_mxfp8_grad(
self,
*,
group_size: int = 4,
weight_shape: tuple[int, int] = (256, 256),
split_alignment: int = 128,
dtype: torch.dtype = torch.bfloat16,
device: torch.device = "cuda",
bias: bool,
) -> None:
"""Rowwise-only MXFP8 GroupedTensor grad output (FP8 token dispatch, backward).

Mirrors ``test_grouped_linear_prequantized_mxfp8_input`` on the backward
side: the rowwise data feeds the dgrad GEMM and the columnwise copy is
manufactured for wgrad. With ``bias`` the per-group bias gradient rides
the columnwise stage of the quantize kernel.
"""
if not mxfp8_available:
pytest.skip(reason_for_no_mxfp8)
maybe_skip_quantization("mxfp8", dims=weight_shape, device=device, dtype=dtype)

# Split sizes (including an empty group)
split_sizes = [split_alignment * i for i in range(group_size)]
random.shuffle(split_sizes)
split_sizes = torch.tensor(split_sizes, dtype=torch.int, device=device)

out_features, in_features = weight_shape
total_tokens = int(split_sizes.sum().item())

x = torch.rand((total_tokens, in_features), dtype=dtype, device=device) - 0.5
dy_hp = torch.rand((total_tokens, out_features), dtype=dtype, device=device) - 0.5

# Wire-format grad output and its exact dequantization (the reference grad).
dy_wire = self._make_rowwise_mxfp8_wire_input(dy_hp, group_size, split_sizes)
dy_ref = tex.group_dequantize(dy_wire, TE_DType[dtype]).rowwise_data.view(
total_tokens, out_features
)

recipe = make_recipe("mxfp8")
op = te.ops.GroupedLinear(
group_size, in_features, out_features, bias=bias, device=device, dtype=dtype
)

def _run(grad):
x_in = x.detach().clone().requires_grad_()
with te.autocast(enabled=True, recipe=recipe):
y = op(x_in, split_sizes)
# Deliver ``grad`` as the op's grad output, as FP8 dispatch would.
_InjectGrad.apply(y, grad).backward(torch.ones_like(y))
wgrads, bgrads = [], []
for group_idx in range(group_size):
weight = getattr(op, f"weight{group_idx}")
wgrads.append(weight.grad.detach().clone())
weight.grad = None
if bias:
bias_param = getattr(op, f"bias{group_idx}")
bgrads.append(bias_param.grad.detach().clone())
bias_param.grad = None
return x_in.grad, wgrads, bgrads

dx_ref, wgrads_ref, bgrads_ref = _run(dy_ref)
dx_test, wgrads_test, bgrads_test = _run(dy_wire)

# Bit-exact match expected (identical quantized grads and kernels).
torch.testing.assert_close(dx_test, dx_ref, rtol=0, atol=0)
for wgrad_test, wgrad_ref in zip(wgrads_test, wgrads_ref):
torch.testing.assert_close(wgrad_test, wgrad_ref, rtol=0, atol=0)
for bgrad_test, bgrad_ref in zip(bgrads_test, bgrads_ref):
torch.testing.assert_close(bgrad_test, bgrad_ref, rtol=0, atol=0)

@pytest.mark.parametrize("dtype", (torch.bfloat16, torch.float16))
@pytest.mark.parametrize(
"quantization",
Expand Down Expand Up @@ -1076,6 +1252,193 @@ def _make_module():
assert_close(fc1.weight.grad, fc1_w_ref_grad, **tols)
assert_close(fc2.weight.grad, fc2_w_ref_grad, **tols)

@pytest.mark.parametrize("weight_requires_grad", (False, True))
def test_grouped_mlp_prequantized_mxfp8_input(
self,
*,
group_size: int = 4,
hidden_size: int = 256,
split_alignment: int = 256,
dtype: torch.dtype = torch.bfloat16,
device: torch.device = "cuda",
weight_requires_grad: bool,
) -> None:
"""Fused grouped MLP with a rowwise-only MXFP8 GroupedTensor input.

Production (FP8 token dispatch) path: FC1 receives an already
rowwise-quantized input with compact scales. The fused op must feed the
rowwise data to the forward GEMM and manufacture FC1's columnwise copy
for wgrad. Compared bit-for-bit against a run on the dequantized wire
input (see ``test_grouped_linear_prequantized_mxfp8_input``).
"""
if not mxfp8_available:
pytest.skip(reason_for_no_mxfp8)
if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported():
pytest.skip("Fused grouped MLP (CuTeDSL) is not supported on this system")
maybe_skip_quantization(
"mxfp8", dims=(hidden_size, hidden_size), device=device, dtype=dtype
)

# Split sizes (including an empty group); sum is a multiple of 128.
split_sizes = [split_alignment * i for i in range(group_size)]
random.shuffle(split_sizes)
split_sizes = torch.tensor(split_sizes, dtype=torch.int, device=device)
total_tokens = int(split_sizes.sum().item())
glu_interleave_size = 32

# Wire-format FC1 input and its exact dequantization (reference input).
x_hp = torch.rand((total_tokens, hidden_size), dtype=dtype, device=device) - 0.5
x_wire = TestGroupedLinearOp._make_rowwise_mxfp8_wire_input(x_hp, group_size, split_sizes)
x_ref = tex.group_dequantize(x_wire, TE_DType[dtype]).rowwise_data.view(
total_tokens, hidden_size
)

probs = torch.rand((total_tokens,), dtype=dtype, device=device)
dy = torch.rand((total_tokens, hidden_size), dtype=dtype, device=device) - 0.5

recipe = make_recipe("mxfp8")
with te.quantized_model_init(enabled=True, recipe=recipe):
fc1 = te.ops.GroupedLinear(
group_size, hidden_size, 2 * hidden_size, bias=False, device=device, dtype=dtype
)
fc2 = te.ops.GroupedLinear(
group_size, hidden_size, hidden_size, bias=False, device=device, dtype=dtype
)
module = te.ops.Sequential(
fc1, te.ops.ScaledSwiGLU(glu_interleave_size=glu_interleave_size), fc2
)
with torch.no_grad():
for param in module.parameters():
param.requires_grad_(requires_grad=weight_requires_grad)

def _run(x):
with te.autocast(enabled=True, recipe=recipe):
y = module(x, split_sizes, probs, split_sizes)
fc1_wgrads, fc2_wgrads = [], []
if weight_requires_grad:
y.backward(dy)
for group_idx in range(group_size):
fc1_w = getattr(fc1, f"weight{group_idx}")
fc2_w = getattr(fc2, f"weight{group_idx}")
fc1_wgrads.append(fc1_w.grad.detach().clone())
fc2_wgrads.append(fc2_w.grad.detach().clone())
fc1_w.grad = None
fc2_w.grad = None
return y.detach(), fc1_wgrads, fc2_wgrads

y_ref, fc1_wgrads_ref, fc2_wgrads_ref = _run(x_ref)
y_test, fc1_wgrads_test, fc2_wgrads_test = _run(x_wire)

# Confirm the CuTeDSL fused op was actually formed (not the fallback).
forward_ops = module._module_groups[0]._forward_ops
assert len(forward_ops) == 1
assert isinstance(forward_ops[0][0], te.ops.fused.GroupedMLP_CuTeGEMMGLU)

# Bit-exact match expected (identical quantized inputs and kernels).
torch.testing.assert_close(y_test, y_ref, rtol=0, atol=0)
for wgrad_test, wgrad_ref in zip(fc1_wgrads_test, fc1_wgrads_ref):
torch.testing.assert_close(wgrad_test, wgrad_ref, rtol=0, atol=0)
for wgrad_test, wgrad_ref in zip(fc2_wgrads_test, fc2_wgrads_ref):
torch.testing.assert_close(wgrad_test, wgrad_ref, rtol=0, atol=0)

@pytest.mark.parametrize("bias", (False, True))
def test_grouped_mlp_prequantized_mxfp8_grad(
self,
*,
group_size: int = 4,
hidden_size: int = 256,
split_alignment: int = 256,
dtype: torch.dtype = torch.bfloat16,
device: torch.device = "cuda",
bias: bool,
) -> None:
"""Fused grouped MLP with a rowwise-only MXFP8 GroupedTensor grad output.

FC2 receives the pre-quantized grad, as an FP8 token dispatch delivers it
on the backward pass. With ``bias`` FC2 uses ``scale_bias``, whose
dbias/dscales need the dequantized grad rather than the fused dbias.
"""
if not mxfp8_available:
pytest.skip(reason_for_no_mxfp8)
if not te.ops.fused.GroupedMLP_CuTeGEMMGLU.is_supported():
pytest.skip("Fused grouped MLP (CuTeDSL) is not supported on this system")
maybe_skip_quantization(
"mxfp8", dims=(hidden_size, hidden_size), device=device, dtype=dtype
)

# Split sizes (including an empty group); sum is a multiple of 128.
split_sizes = [split_alignment * i for i in range(group_size)]
random.shuffle(split_sizes)
split_sizes = torch.tensor(split_sizes, dtype=torch.int, device=device)
total_tokens = int(split_sizes.sum().item())
glu_interleave_size = 32

x = torch.rand((total_tokens, hidden_size), dtype=dtype, device=device) - 0.5
probs = torch.rand((total_tokens,), dtype=dtype, device=device)
dy_hp = torch.rand((total_tokens, hidden_size), dtype=dtype, device=device) - 0.5

# Wire-format grad output and its exact dequantization (the reference grad).
dy_wire = TestGroupedLinearOp._make_rowwise_mxfp8_wire_input(dy_hp, group_size, split_sizes)
dy_ref = tex.group_dequantize(dy_wire, TE_DType[dtype]).rowwise_data.view(
total_tokens, hidden_size
)

recipe = make_recipe("mxfp8")
with te.quantized_model_init(enabled=True, recipe=recipe):
fc1 = te.ops.GroupedLinear(
group_size, hidden_size, 2 * hidden_size, bias=bias, device=device, dtype=dtype
)
fc2 = te.ops.GroupedLinear(
group_size,
hidden_size,
hidden_size,
bias=bias,
device=device,
dtype=dtype,
scale_bias=bias,
)
module = te.ops.Sequential(
fc1, te.ops.ScaledSwiGLU(glu_interleave_size=glu_interleave_size), fc2
)

def _run(grad):
x_in = x.detach().clone().requires_grad_()
fc2_extra = (split_sizes, probs) if bias else (split_sizes,)
with te.autocast(enabled=True, recipe=recipe):
y = module(x_in, split_sizes, probs, *fc2_extra)
_InjectGrad.apply(y, grad).backward(torch.ones_like(y))
grads = [("dx", x_in.grad)]
for name, fc in (("fc1", fc1), ("fc2", fc2)):
for group_idx in range(group_size):
weight = getattr(fc, f"weight{group_idx}")
grads.append((f"{name}_w{group_idx}", weight.grad.detach().clone()))
weight.grad = None
if bias:
bias_param = getattr(fc, f"bias{group_idx}")
grads.append((f"{name}_b{group_idx}", bias_param.grad.detach().clone()))
bias_param.grad = None
return grads

grads_ref = _run(dy_ref)
grads_test = _run(dy_wire)

# Confirm the CuTeDSL fused op was actually formed (not the fallback).
forward_ops = module._module_groups[0]._forward_ops
assert len(forward_ops) == 1
assert isinstance(forward_ops[0][0], te.ops.fused.GroupedMLP_CuTeGEMMGLU)

# Bit-exact match expected (identical quantized grads and kernels), except
# bias gradients: the fused kernels generate them with an accumulation
# that is not reproducible run to run (two runs on identical inputs differ
# by one BF16 ulp). Same tolerances as
# ``test_grouped_mlp_single_weight_numerics``.
bias_tols = {"rtol": 0.05, "atol": 0.015625}
for (name, grad_test), (_, grad_ref) in zip(grads_test, grads_ref):
if "_b" in name:
torch.testing.assert_close(grad_test, grad_ref, **bias_tols)
else:
torch.testing.assert_close(grad_test, grad_ref, rtol=0, atol=0)

@pytest.mark.parametrize("bias", (False, True))
@pytest.mark.parametrize("quantization", _grouped_mlp_quantization_list)
@pytest.mark.parametrize(
Expand Down
Loading