Hacktoberfest 2026: những issue maintainer đã đánh dấu cho tháng Mười, đang mở và phù hợp người mới. Xem issue Hacktoberfest

Enabling grouped MLP fusion breaks previously working 128-aligned inputs

Đang mở
#3,640 8 bình luận 0 reaction 0 người được giao Xem trên GitHub

Maintainer thường phản hồi trong vòng 2 ngày

Chưa có ai nhận issue này.

Đánh giá

Độ khó
4/5
Thời gian dự kiến
3-5 ngày
Mức phù hợp với người mới
54/100
Loại issue
Lỗi
Độ rõ ràng
Đặc tả rõ ràng
Mức độ hoạt động
Sôi nổi
Công nghệ
python, pytorch
Lĩnh vực
machine-learning

Hướng nghiên cứu

Save the supplied reproducer as repro_alignment.py and run it with NVTE_CUTEDSL_FUSED_GROUPED_MLP=1; compare the failing 128-row case with the padded and fusion-disabled runs. Start by tracing how grouped MLP fusion is selected and where per-expert row counts are checked. Done means misaligned groups no longer launch the failing fused kernel, and regression coverage verifies the fallback or other chosen handling.

Do mô hình lập chỉ mục viết ra từ nội dung của issue.

Mô tả

bug

My agent says:

The same MXFP8 grouped MLP inputs with 128-aligned expert sizes work correctly without fusion, but produce incorrect results when fusion is enabled. The fused cuDNN backend requires each expert’s row count to be a multiple of 256, while TE checks only the total row count and still selects the fusion.

On B100, 128 rows/expert produce incorrect, nondeterministic outputs and gradients; 384 rows/expert can cause illegal memory access. Padding each expert to 256 or disabling fusion passes the reproducer.

Tests missed this because the main fused numerical tests already generate 256-aligned groups. TE should enforce the requirement, pad, or fall back before launching the kernel.

Environment: B100, TE 2.21.0.dev0, PyTorch 2.15.0a0, CUDA 13.4, cuDNN frontend 1.29.0, CUTLASS DSL 4.8.0.dev0.

Tested reproducer (expand)

Save as repro_alignment.py; run each command in a fresh process:

export CUDA_VISIBLE_DEVICES=0 NVTE_FLASH_ATTN_V4=0 OMP_NUM_THREADS=8
NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 python repro_alignment.py               # fails
NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 python repro_alignment.py --pad-to 256  # passes
NVTE_CUTEDSL_FUSED_GROUPED_MLP=0 python repro_alignment.py               # passes
import argparse
import importlib.metadata
import os

os.environ.setdefault("NVTE_CUTEDSL_FUSED_GROUPED_MLP", "1")
os.environ["NVTE_GROUPED_LINEAR_SINGLE_PARAM"] = "0"
os.environ["NVTE_ALLOW_NONDETERMINISTIC_ALGO"] = "1"

import torch
import transformer_engine
import transformer_engine.pytorch as te
from transformer_engine.common.recipe import MXFP8BlockScaling

parser = argparse.ArgumentParser()
parser.add_argument("--rows", type=int, default=128)
parser.add_argument("--pad-to", type=int, choices=(128, 256), default=128)
args = parser.parse_args()
assert torch.cuda.is_available()
print("torch:", torch.__version__, "CUDA:", torch.version.cuda,
      "TE:", transformer_engine.__version__,
      "cuDNN frontend:", importlib.metadata.version("nvidia-cudnn-frontend"),
      "CUTLASS DSL:", importlib.metadata.version("nvidia-cutlass-dsl"),
      "GPU:", torch.cuda.get_device_name(), flush=True)

groups, hidden, ffn = 8, 2048, 1024
physical_rows = (args.rows + args.pad_to - 1) // args.pad_to * args.pad_to
torch.manual_seed(0)
common = dict(bias=False, dtype=torch.bfloat16, device="cuda")
model = te.ops.Sequential(
    te.ops.GroupedLinear(groups, hidden, 2 * ffn, **common),
    te.ops.ScaledSwiGLU(glu_interleave_size=32),
    te.ops.GroupedLinear(groups, ffn, hidden, **common),
)
torch.manual_seed(2)
data = torch.randn(groups, args.rows, hidden, dtype=torch.bfloat16, device="cuda")
probabilities = torch.rand(groups, args.rows, dtype=torch.bfloat16, device="cuda")
splits = torch.full((groups,), physical_rows, dtype=torch.int64, device="cuda")
baseline = None

for iteration in range(5):
    x = data.clone().requires_grad_()
    probs = probabilities.clone().requires_grad_()
    padded_x = torch.nn.functional.pad(x, (0, 0, 0, physical_rows - args.rows))
    padded_probs = torch.nn.functional.pad(probs, (0, physical_rows - args.rows))
    model.zero_grad(set_to_none=True)
    with te.autocast(enabled=True, recipe=MXFP8BlockScaling()):
        output = model(padded_x.reshape(-1, hidden), splits, padded_probs.reshape(-1), splits)
    output = output.reshape(groups, physical_rows, hidden)[:, :args.rows]
    output.backward(torch.ones_like(output))
    torch.cuda.synchronize()

    if iteration == 0:
        selected = [type(op[0]).__name__ for op in model._module_groups[0]._forward_ops]
        print("Selected ops:", selected, "Rows per expert:", physical_rows, flush=True)
        expected_fused = os.environ["NVTE_CUTEDSL_FUSED_GROUPED_MLP"] == "1"
        assert ("GroupedMLP_CuTeGEMMGLU" in selected) == expected_fused

    current = {"output": output.detach(), "dx": x.grad, "dprob": probs.grad}
    current.update({name: parameter.grad for name, parameter in model.named_parameters()})
    assert all(value is not None and torch.isfinite(value).all() for value in current.values())
    if baseline is None:
        baseline = {name: value.clone() for name, value in current.items()}
    else:
        differences = {
            name: (value.float() - baseline[name].float()).abs().max().item()
            for name, value in current.items() if not torch.equal(value, baseline[name])
        }
        print("Repeat", iteration, "max absolute differences:", differences, flush=True)
        assert not differences, "Identical inputs/weights produced different outputs or gradients"
print("PASS: all five forward/backward passes are bitwise identical", flush=True)
Ngôn ngữ chính
Python
Star
3.6k
Fork
851
Merge trung bình
4 ngày 15 giờ
Pull request đã merge (30 ngày)
51

Chuẩn bị môi trường

Bắt đầu từ đâu

  1. Đọc hết issue, rồi đọc hướng dẫn đóng góp của dự án.
  2. Bình luận trên issue rằng bạn sẽ nhận — tránh hai người làm cùng một việc.
  3. Fork repository và làm thay đổi trên một nhánh.
  4. Mở pull request có tham chiếu số hiệu của issue.

Issue khác của NVIDIA/TransformerEngine

Tất cả issue của NVIDIA/TransformerEngine

Issue tương tự

Thêm issue về Python

Nhận issue mới trong hộp thư của bạn

Bản tóm tắt ngắn những issue GitHub phù hợp với người mới.