Llama4 MoE Grouped GEMM (#2639)

* add llama4 reference layer

* add llama4 reference impl

* formatting
This commit is contained in:
jeromeku 2025-05-28 03:26:35 -07:00 committed by GitHub
commit 15bce315ac
9 changed files with 1418 additions and 372 deletions

View file

@ -43,18 +43,26 @@ sets of tests. E.g., to run forward tests with autotune turned on: `pytest -sv
- `grouped_gemm/tests/test_qwen3_moe.py`: end to end test for Qwen3 MoE block. IMPORTANT: read `tests/run_qwen3_moe_tests.sh` as well as notes in the test itself for complications when running parametrized pytest test suites and triton / autotune. TLDR: use the test script and NOT pytest to run the tests.
### Benchmarks
- `grouped_gemm/benchmark/benchmark_fused_moe.py`: benchmarks HF `Qwen3SpareMOEBlock` against the fused implementation
- `grouped_gemm/benchmark/benchmark_fused_moe.py`: benchmarks HF `Qwen3SpareMOEBlock` or `Llama4TextMoe` against the fused implementation
Running with these flags on an `H100` to bench forward pass (run with `--help` to see all available flags):
For `Qwen3-30B-A3B`:
```
python benchmark/benchmark_fused_moe.py --mode forward --seqlen 1024 --permute_x --permute_y --autotune
python benchmark/benchmark_fused_moe.py --model qwen3 --mode forward --seqlen 1024 --permute_x --permute_y --autotune
```
For the backward bench:
```
python benchmark/benchmark_fused_moe.py --mode backward --seqlen 1024 --permute_x --permute_y --autotune
python benchmark/benchmark_fused_moe.py --model qwen3 --mode backward --seqlen 1024 --permute_x --permute_y --autotune
```
On my machine and env, I get speedups > 25x and 14x respectively.
For `Llama-4-Scout-17B-16E`:
```
python benchmark/benchmark_fused_moe.py --model llama4 --autotune --mode=forward --permute_y
```
Ditto for backwards.
### Notes
- Tested and benched on `H100`, though should run on Ampere and possibly even earlier gpu generations though the autotuning configs will need to be adjusted.
@ -62,6 +70,7 @@ On my machine and env, I get speedups > 25x and 14x respectively.
- The kernels can be run either as autotuned (see `autotuning.py`) or with manually specified config (see `tuning.py`). Recommended to run using autotuner since the MoE block requires 2 configs for the forward (2 grouped gemms) and 4 for the backwards (dX and dW per grouped gemm, 2 grouped gemms).
- Running with autotuning turned off with the default manual kernel config will result is **highly** sub-optimal performance as it is only meant for testing / debugging purposes.
- I've tried to strike a balance between compilation time and autotuning search space -- can probably squeeze even more performance for specific workloads.
- The Llama4 reference layer is still highly under-optimized as there are many low-hanging opportunities for further speedups around routing and shared expert calculation.
TODO:
- TMA store: implemented but not enabled currently due to non-determinism arising from triton pipelining bug.
@ -69,4 +78,8 @@ TODO:
- Additional optimizations:
- Fused / optimized implementations of routing, token sorting, etc.
- Better software pipelining within grouped gemm
- Threadblock swizzling for better L2 caching
- Threadblock swizzling for better L2 caching
- Llama4
- Fused gather / topk weight merging
- Custom topk, gather indices kernel
- Shared expert fusion with experts calculation

View file

@ -1,7 +1,22 @@
import argparse
import time
from contextlib import nullcontext
import torch
from transformers import AutoConfig
from transformers.models.llama4 import Llama4TextConfig
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
from transformers.models.qwen3_moe import Qwen3MoeConfig
from transformers.models.qwen3_moe.modeling_qwen3_moe import Qwen3MoeSparseMoeBlock
from triton.testing import do_bench
from utils import (
create_kernel_configs,
get_autotuner,
post_process_results,
postprocess_autotune_results,
save_results,
)
from grouped_gemm.kernels.autotuning import (
DEFAULT_K_BLOCK_SIZES,
DEFAULT_M_BLOCK_SIZES,
@ -16,60 +31,50 @@ from grouped_gemm.kernels.tuning import (
KernelResult,
TritonTuningContext,
)
from grouped_gemm.reference.moe_block import Qwen3MoeFusedGroupedGEMMBlock
from transformers import AutoConfig
from transformers.models.qwen3_moe import Qwen3MoeConfig
from transformers.models.qwen3_moe.modeling_qwen3_moe import Qwen3MoeSparseMoeBlock
from triton.testing import do_bench
from utils import (
create_kernel_configs,
post_process_results,
save_results,
)
from grouped_gemm.reference.layers.llama4_moe import Llama4TritonTextMoe
from grouped_gemm.reference.layers.qwen3_moe import Qwen3MoeFusedGroupedGEMMBlock
SEED = 42
LLAMA4_ID = "meta-llama/Llama-4-Scout-17B-16E"
QWEN3_MODEL_ID = "Qwen/Qwen3-30B-A3B"
def run_benchmark_forward(
config: Qwen3MoeConfig,
ref_model: torch.nn.Module,
tt_model: torch.nn.Module,
config: AutoConfig,
seqlen: int,
dtype: torch.dtype,
permute_x: bool,
permute_y: bool,
autotune: bool,
kernel_config_fwd: KernelConfigForward = None,
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
bs: int = 1,
):
torch.manual_seed(SEED) # Should not be needed when running using pytest -- autouse fixture in conftest.py
torch.manual_seed(
SEED
) # Should not be needed when running using pytest -- autouse fixture in conftest.py
device = "cuda"
hidden_size = config.hidden_size
bs = 1
# Reference op -- HF
moe_block = Qwen3MoeSparseMoeBlock(config).to(device, dtype)
X = torch.randn(
bs, seqlen, hidden_size, dtype=dtype, device=device, requires_grad=True
)
# Triton kernel grouped gemm version of MoE Block -- this is what we're testing
fused_gemm_block = Qwen3MoeFusedGroupedGEMMBlock.from_hf(
moe_block,
permute_x=permute_x,
permute_y=permute_y,
autotune=autotune,
kernel_config_fwd=kernel_config_fwd,
kernel_config_bwd_dW=kernel_config_bwd_dW,
kernel_config_bwd_dX=kernel_config_bwd_dX,
).to(device, dtype)
X = torch.randn(bs, seqlen, hidden_size, dtype=dtype, device=device, requires_grad=True)
ref_output, _ = moe_block(X)
# Forward
bench_forward_ref = lambda: moe_block(X)
bench_forward_fused = lambda: fused_gemm_block(X)
bench_forward_ref = lambda: ref_model(X) # noqa: E731
bench_forward_fused = lambda: tt_model(X) # noqa: E731
ref_forward_time = do_bench(bench_forward_ref)
with TritonTuningContext(kernel_config_fwd) as ctx:
if not autotune:
assert kernel_config_fwd is not None
tuning_context = TritonTuningContext(kernel_config_fwd)
else:
tuning_context = nullcontext()
with tuning_context:
fused_forward_time = do_bench(bench_forward_fused)
if not ctx.success:
if (not autotune) and (not tuning_context.success):
return 0, 1
print(
@ -77,8 +82,107 @@ def run_benchmark_forward(
)
return ref_forward_time, fused_forward_time
def run_benchmark_backward(
config: Qwen3MoeConfig,
ref_model: torch.nn.Module,
tt_model: torch.nn.Module,
config: AutoConfig,
seqlen: int,
dtype: torch.dtype,
bs=1,
):
torch.manual_seed(
SEED
) # Should not be needed when running using pytest -- autouse fixture in conftest.py
device = "cuda"
hidden_size = config.hidden_size
X = torch.randn(
bs, seqlen, hidden_size, dtype=dtype, device=device, requires_grad=True
)
X_test = X.detach().clone().requires_grad_(True)
output, _ = ref_model(X)
# Prevent autotuning forward pass
from grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
_autotuned_grouped_gemm_forward_kernel.configs = (
_autotuned_grouped_gemm_forward_kernel.configs[:20]
)
test_output, _ = tt_model(X_test)
# Bench
grad_output = torch.randn_like(output)
bench_backward_ref = lambda: output.backward(grad_output, retain_graph=True) # noqa: E731
bench_backward_fused = lambda: test_output.backward(grad_output, retain_graph=True) # noqa: E731
ref_backward_time = do_bench(
bench_backward_ref, grad_to_none=[X, *ref_model.parameters()]
)
fused_backward_time = do_bench(
bench_backward_fused, grad_to_none=[X_test, *tt_model.parameters()]
)
print(
f"Backward: ref {ref_backward_time:.4f}, fused {fused_backward_time:.4f}, speedup {ref_backward_time / fused_backward_time:.1f}x"
)
return ref_backward_time, fused_backward_time
def setup_model(
config: Qwen3MoeConfig | Llama4TextConfig,
dtype,
permute_x,
permute_y,
autotune,
kernel_config_fwd,
kernel_config_bwd_dW,
kernel_config_bwd_dX,
dX_only=False,
dW_only=False,
overlap_router_shared=False,
device="cuda",
):
if isinstance(config, Qwen3MoeConfig):
ref_model = Qwen3MoeSparseMoeBlock(config).to(device, dtype)
# Triton kernel grouped gemm version of MoE Block -- this is what we're testing
tt_model = Qwen3MoeFusedGroupedGEMMBlock.from_hf(
ref_model,
permute_x=permute_x,
permute_y=permute_y,
autotune=autotune,
kernel_config_fwd=kernel_config_fwd,
kernel_config_bwd_dW=kernel_config_bwd_dW,
kernel_config_bwd_dX=kernel_config_bwd_dX,
dX_only=dX_only,
dW_only=dW_only,
).to(device, dtype)
elif isinstance(config, Llama4TextConfig):
ref_model = Llama4TextMoe(config).to(device, dtype)
tt_model = Llama4TritonTextMoe(
config,
overlap_router_shared=overlap_router_shared,
permute_x=permute_x,
permute_y=permute_y,
autotune=autotune,
kernel_config_fwd=kernel_config_fwd,
kernel_config_bwd_dW=kernel_config_bwd_dW,
kernel_config_bwd_dX=kernel_config_bwd_dX,
dX_only=dX_only,
dW_only=dW_only,
).to(device, dtype)
else:
raise ValueError(f"Unrecognized config {type(config).__name__}")
return ref_model, tt_model
def run_benchmark(
mode: str,
model_config: Qwen3MoeConfig | Llama4TextConfig,
seqlen: int,
dtype: torch.dtype,
permute_x: bool,
@ -87,20 +191,21 @@ def run_benchmark_backward(
kernel_config_fwd: KernelConfigForward = None,
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
dX_only: bool = False,
dW_only: bool = False,
overlap_router_shared: bool = False,
results_dir: str = None,
):
torch.manual_seed(SEED) # Should not be needed when running using pytest -- autouse fixture in conftest.py
device = "cuda"
hidden_size = config.hidden_size
bs = 1
if autotune:
autotuner = get_autotuner(mode)
if mode == "dW":
dW_only = True
elif mode == "dX":
dX_only = True
else:
dW_only = dX_only = False
# Reference op -- HF
moe_block = Qwen3MoeSparseMoeBlock(config).to(device, dtype)
# Triton kernel grouped gemm version of MoE Block -- this is what we're testing
fused_gemm_block = Qwen3MoeFusedGroupedGEMMBlock.from_hf(
moe_block,
ref_model, tt_model = setup_model(
model_config,
dtype=dtype,
permute_x=permute_x,
permute_y=permute_y,
autotune=autotune,
@ -109,130 +214,109 @@ def run_benchmark_backward(
kernel_config_bwd_dX=kernel_config_bwd_dX,
dX_only=dX_only,
dW_only=dW_only,
).to(device, dtype)
X = torch.randn(bs, seqlen, hidden_size, dtype=dtype, device=device, requires_grad=True)
X_test = X.detach().clone().requires_grad_(True)
output, _ = moe_block(X)
# Prevent autotuning forward pass
from grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
_autotuned_grouped_gemm_forward_kernel.configs = _autotuned_grouped_gemm_forward_kernel.configs[:20]
test_output, _ = fused_gemm_block(X_test)
# Bench
grad_output = torch.randn_like(output)
bench_backward_ref = lambda: output.backward(grad_output, retain_graph=True) # noqa: E731
bench_backward_fused = lambda: test_output.backward(grad_output, retain_graph=True) # noqa: E731
ref_backward_time = do_bench(bench_backward_ref, grad_to_none=[X, *moe_block.parameters()])
fused_backward_time = do_bench(bench_backward_fused, grad_to_none=[X_test, *fused_gemm_block.parameters()])
print(
f"Backward: ref {ref_backward_time:.4f}, fused {fused_backward_time:.4f}, speedup {ref_backward_time / fused_backward_time:.1f}x"
overlap_router_shared=overlap_router_shared,
)
return ref_backward_time, fused_backward_time
def run_benchmark(
mode: str,
model_config: Qwen3MoeConfig,
seqlen: int,
dtype: torch.dtype,
permute_x: bool,
permute_y: bool,
autotune: bool,
kernel_config_fwd: KernelConfigForward = None,
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
):
if mode == "forward":
ref_time, fused_time = run_benchmark_forward(
model_config,
seqlen,
dtype,
permute_x,
permute_y,
autotune,
kernel_config_fwd,
kernel_config_bwd_dW,
kernel_config_bwd_dX,
ref_model,
tt_model,
config=model_config,
seqlen=seqlen,
dtype=dtype,
autotune=autotune,
kernel_config_fwd=kernel_config_fwd,
)
elif mode == "dW":
else:
ref_time, fused_time = run_benchmark_backward(
model_config,
seqlen,
dtype,
permute_x,
permute_y,
autotune,
kernel_config_fwd,
kernel_config_bwd_dW,
kernel_config_bwd_dX,
dW_only=True,
)
elif mode == "dX":
ref_time, fused_time = run_benchmark_backward(
model_config,
seqlen,
dtype,
permute_x,
permute_y,
autotune,
kernel_config_fwd,
kernel_config_bwd_dW,
kernel_config_bwd_dX,
dX_only=True,
)
elif mode == "backward":
ref_time, fused_time = run_benchmark_backward(
model_config,
seqlen,
dtype,
permute_x,
permute_y,
autotune,
kernel_config_fwd,
kernel_config_bwd_dW,
kernel_config_bwd_dX,
dX_only=False,
dW_only=False,
ref_model, tt_model, config=model_config, seqlen=seqlen, dtype=dtype
)
if autotune:
if mode == "backward":
autotuner_dW, autotuner_dX = autotuner
postprocess_autotune_results(
autotuner_dW, "dW", ref_time, fused_time, results_dir
)
postprocess_autotune_results(
autotuner_dX, "dX", ref_time, fused_time, results_dir
)
else:
postprocess_autotune_results(
autotuner, mode, ref_time, fused_time, results_dir
)
return ref_time, fused_time
# NOTE: better to use autotuner for now, since the MoE block needs 2 different kernel configs for forward (2 grouped gemms, gate_up_proj and down_proj)
# and the backward pass needs 4 different kernel configs (2 grouped gemms each for dW and dX)
# The benchmark only supports 1 kernel config at a time so the same config will be used for both grouped gemms, which is suboptimal.
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--results_dir", type=str, default="benchmark_results")
parser.add_argument("--model", type=str, choices=["llama4", "qwen3"], required=True)
parser.add_argument("--seqlen", type=int, default=1024)
parser.add_argument("--dtype", type=str, choices=["bfloat16", "float16"], default="bfloat16")
parser.add_argument(
"--dtype", type=str, choices=["bfloat16", "float16"], default="bfloat16"
)
parser.add_argument("--permute_x", action="store_true")
parser.add_argument("--permute_y", action="store_true")
parser.add_argument("--autotune", action="store_true")
parser.add_argument("--BLOCK_SIZE_M", nargs=2, type=int, default=[DEFAULT_M_BLOCK_SIZES[0], DEFAULT_M_BLOCK_SIZES[-1]])
parser.add_argument("--BLOCK_SIZE_N", nargs=2, type=int, default=[DEFAULT_N_BLOCK_SIZES[0], DEFAULT_N_BLOCK_SIZES[-1]])
parser.add_argument("--BLOCK_SIZE_K", nargs=2, type=int, default=[DEFAULT_K_BLOCK_SIZES[0], DEFAULT_K_BLOCK_SIZES[-1]])
parser.add_argument("--num_warps", nargs=2, type=int, default=[DEFAULT_NUM_WARPS[0], DEFAULT_NUM_WARPS[-1]])
parser.add_argument("--num_stages", nargs=2, type=int, default=[DEFAULT_NUM_STAGES[0], DEFAULT_NUM_STAGES[-1]])
parser.add_argument("--use_tma_load_w", action="store_true") # No need to specify, will automatically parametrize these for each kernel config
parser.add_argument("--use_tma_load_x", action="store_true") # No need to specify, will automatically parametrize these for each kernel config
parser.add_argument("--use_tma_load_dy", action="store_true") # No need to specify, will automatically parametrize these for each kernel config
parser.add_argument("--mode", type=str, choices=["forward", "backward", "dW", "dX"], default="forward")
parser.add_argument("--overlap_router_shared", action="store_true")
parser.add_argument(
"--BLOCK_SIZE_M",
nargs=2,
type=int,
default=[DEFAULT_M_BLOCK_SIZES[0], DEFAULT_M_BLOCK_SIZES[-1]],
)
parser.add_argument(
"--BLOCK_SIZE_N",
nargs=2,
type=int,
default=[DEFAULT_N_BLOCK_SIZES[0], DEFAULT_N_BLOCK_SIZES[-1]],
)
parser.add_argument(
"--BLOCK_SIZE_K",
nargs=2,
type=int,
default=[DEFAULT_K_BLOCK_SIZES[0], DEFAULT_K_BLOCK_SIZES[-1]],
)
parser.add_argument(
"--num_warps",
nargs=2,
type=int,
default=[DEFAULT_NUM_WARPS[0], DEFAULT_NUM_WARPS[-1]],
)
parser.add_argument(
"--num_stages",
nargs=2,
type=int,
default=[DEFAULT_NUM_STAGES[0], DEFAULT_NUM_STAGES[-1]],
)
parser.add_argument(
"--use_tma_load_w", action="store_true"
) # No need to specify, will automatically parametrize these for each kernel config
parser.add_argument(
"--use_tma_load_x", action="store_true"
) # No need to specify, will automatically parametrize these for each kernel config
parser.add_argument(
"--use_tma_load_dy", action="store_true"
) # No need to specify, will automatically parametrize these for each kernel config
parser.add_argument(
"--mode",
type=str,
choices=["forward", "backward", "dW", "dX"],
default="forward",
)
args = parser.parse_args()
args.dtype = getattr(torch, args.dtype)
model_id = "Qwen/Qwen3-30B-A3B"
model_id = QWEN3_MODEL_ID if args.model == "qwen3" else LLAMA4_ID
model_config = AutoConfig.from_pretrained(model_id)
model_config = model_config.text_config if args.model == "llama4" else model_config
mode = args.mode
if args.autotune:
# logging.basicConfig(level=logging.INFO)
print(
f"Benchmarking {model_id} {mode}: seqlen={args.seqlen}, dtype={args.dtype}, permute_x={args.permute_x}, permute_y={args.permute_y}, autotune"
)
@ -245,16 +329,28 @@ if __name__ == "__main__":
permute_x=args.permute_x,
permute_y=args.permute_y,
autotune=args.autotune,
overlap_router_shared=args.overlap_router_shared,
results_dir=args.results_dir,
)
end_time = time.time()
print(f"Total time: {end_time - start_time:.4f} seconds")
# NOTE: better to use autotuner for now, since the MoE block needs 2 different kernel configs for forward (2 grouped gemms, gate_up_proj and down_proj)
# and the backward pass needs 4 different kernel configs (2 grouped gemms each for dW and dX)
# The benchmark only supports 1 kernel config at a time so the same config will be used for both grouped gemms, which is suboptimal.
else:
assert False, "Use autotune for now"
kernel_configs = create_kernel_configs(args, args.permute_x, args.permute_y)
print(f"Running {len(kernel_configs)} kernel configs")
default_kernel_config_fwd = KernelConfigForward(permute_x=args.permute_x, permute_y=args.permute_y)
default_kernel_config_bwd_dW = KernelConfigBackward_dW(permute_x=args.permute_x, permute_y=args.permute_y)
default_kernel_config_bwd_dX = KernelConfigBackward_dX(permute_x=args.permute_x, permute_y=args.permute_y)
default_kernel_config_fwd = KernelConfigForward(
permute_x=args.permute_x, permute_y=args.permute_y
)
default_kernel_config_bwd_dW = KernelConfigBackward_dW(
permute_x=args.permute_x, permute_y=args.permute_y
)
default_kernel_config_bwd_dX = KernelConfigBackward_dX(
permute_x=args.permute_x, permute_y=args.permute_y
)
results = []
for kernel_config in kernel_configs:
if args.mode == "forward":
@ -274,7 +370,7 @@ if __name__ == "__main__":
print(
f"Benchmarking {model_id} {args.mode} with seqlen={args.seqlen}, dtype={args.dtype}, permute_x={args.permute_x}, permute_y={args.permute_y}, kernel_config_fwd={kernel_config_fwd}, kernel_config_bwd_dW={kernel_config_bwd_dW}, kernel_config_bwd_dX={kernel_config_bwd_dX}"
)
ref_time, fused_time = run_benchmark(
args.mode,
model_config,
@ -287,11 +383,17 @@ if __name__ == "__main__":
kernel_config_bwd_dW=kernel_config_bwd_dW,
kernel_config_bwd_dX=kernel_config_bwd_dX,
)
results.append(KernelResult(
torch_time=ref_time,
triton_time=fused_time,
speedup=ref_time / fused_time,
kernel_config=kernel_config,
))
df = post_process_results(results, args.mode, args.seqlen, args.dtype, args.autotune)
save_results(df, args.results_dir, args.mode, args.seqlen, args.dtype, args.autotune)
results.append(
KernelResult(
torch_time=ref_time,
triton_time=fused_time,
speedup=ref_time / fused_time,
kernel_config=kernel_config,
)
)
df = post_process_results(
results, args.mode, args.seqlen, args.dtype, args.autotune
)
save_results(
df, args.results_dir, args.mode, args.seqlen, args.dtype, args.autotune
)

View file

@ -8,6 +8,7 @@ from itertools import product
import pandas as pd
import torch
from grouped_gemm.kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
@ -22,7 +23,12 @@ def create_merged_results(
df: pd.DataFrame, mode: str, seqlen: int, dtype: torch.dtype, autotune: bool
):
kernel_result_cols = df.columns.to_list()
test_config_dict = {"mode": mode, "seqlen": seqlen, "dtype": dtype, "autotune": autotune}
test_config_dict = {
"mode": mode,
"seqlen": seqlen,
"dtype": dtype,
"autotune": autotune,
}
test_config_cols = list(test_config_dict.keys())
for col in test_config_cols:
df[col] = test_config_dict[col]
@ -65,12 +71,28 @@ def create_kernel_configs(args: argparse.Namespace, permute_x: bool, permute_y:
block_n_range = power_of_two_range(args.BLOCK_SIZE_N[0], args.BLOCK_SIZE_N[1])
block_k_range = power_of_two_range(args.BLOCK_SIZE_K[0], args.BLOCK_SIZE_K[1])
num_warps_range = multiples_of_range(args.num_warps[0], args.num_warps[1], step=2)
num_stages_range = multiples_of_range(args.num_stages[0], args.num_stages[1], step=1)
num_stages_range = multiples_of_range(
args.num_stages[0], args.num_stages[1], step=1
)
mode = args.mode
kernel_configs = []
for block_m, block_n, block_k, num_warps, num_stages, tma_load_a, tma_load_b in product(
block_m_range, block_n_range, block_k_range, num_warps_range, num_stages_range, [True, False], [True, False]
for (
block_m,
block_n,
block_k,
num_warps,
num_stages,
tma_load_a,
tma_load_b,
) in product(
block_m_range,
block_n_range,
block_k_range,
num_warps_range,
num_stages_range,
[True, False],
[True, False],
):
if mode == "forward":
kernel_config = KernelConfigForward(
@ -141,3 +163,66 @@ def power_of_two_range(start, end):
def multiples_of_range(start, end, step=1):
return list(range(start, end + step, step))
def map_key_to_args(key, mode):
pass
def save_autotune_results(autotune_cache, mode, ref_time, fused_time, results_dir):
device_name = torch.cuda.get_device_name().replace(" ", "_")
dt = datetime.datetime.now().strftime("%Y%m%d_%H%M")
save_dir = f"{results_dir}/{mode}/autotune/{dt}/{device_name}"
if not os.path.exists(save_dir):
os.makedirs(save_dir)
for key, config in autotune_cache.items():
key = [
str(k) if not "torch" in str(k) else str(k.split("torch.")[-1]) for k in key
]
filename = "_".join(key)
save_path = f"{save_dir}/{filename}.json"
print(f"Saving autotune results to {save_path}")
with open(save_path, "w") as f:
result = {
**config.all_kwargs(),
"ref_time": ref_time,
"fused_time": fused_time,
}
json.dump(result, f)
def get_autotuner(mode):
if mode == "forward":
from grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
return _autotuned_grouped_gemm_forward_kernel
elif mode == "dW":
from grouped_gemm.kernels.backward import _autotuned_grouped_gemm_dW_kernel
return _autotuned_grouped_gemm_dW_kernel
elif mode == "dX":
from grouped_gemm.kernels.backward import _autotuned_grouped_gemm_dX_kernel
return _autotuned_grouped_gemm_dX_kernel
elif mode == "backward":
from grouped_gemm.kernels.backward import (
_autotuned_grouped_gemm_dW_kernel,
_autotuned_grouped_gemm_dX_kernel,
)
return _autotuned_grouped_gemm_dW_kernel, _autotuned_grouped_gemm_dX_kernel
else:
raise ValueError(f"Invalid mode: {mode}")
def postprocess_autotune_results(autotuner, mode, ref_time, fused_time, results_dir):
for key, value in autotuner.cache.items():
print(f"{mode} {key}: {value.all_kwargs()}")
save_autotune_results(
autotuner.cache,
mode=mode,
ref_time=ref_time,
fused_time=fused_time,
results_dir=results_dir,
)

View file

@ -0,0 +1,434 @@
from dataclasses import dataclass
from typing import Tuple
import torch
import torch.nn.functional as F
from transformers.models.llama4 import Llama4TextConfig
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
get_routing_indices,
permute,
torch_grouped_gemm,
unpermute,
)
"""
Reference implementation of Llama4 MoE block using triton grouped gemm.
`Llama4GroupedGemmTextMoe` is the HF `Llama4TextMoe` block implemented with a torch-native grouped gemm.
`Llama4TritonTextMoe` is the HF `Llama4TextMoe` implemented with triton grouped gemm.
"""
@dataclass
class Llama4MoeResult:
token_counts_by_expert: torch.Tensor
gather_indices: torch.Tensor
topk_weights: torch.Tensor
hidden_states_after_weight_merge: torch.Tensor
first_gemm: torch.Tensor
intermediate: torch.Tensor
second_gemm: torch.Tensor
hidden_states_unpermute: torch.Tensor
shared_expert_out: torch.Tensor
final_out: torch.Tensor
router_logits: torch.Tensor = None
class Llama4GroupedGemmTextMoe(Llama4TextMoe):
EXPERT_WEIGHT_NAMES = ["experts.gate_up_proj", "experts.down_proj"]
def __init__(
self,
config: Llama4TextConfig,
overlap_router_shared=False,
verbose=False,
debug=False,
):
super().__init__(config)
self.overlap_router_shared = overlap_router_shared
self.verbose = verbose
self.debug = debug
# Permute in-place expert weights
E, K, N = self.num_experts, self.hidden_dim, self.experts.expert_dim
assert self.experts.gate_up_proj.shape == torch.Size([E, K, 2 * N]), (
f"{self.experts.gate_up_proj.shape} != {[E, K, 2 * N]}"
)
permuted_shape = [E, 2 * N, K]
permuted_stride = [2 * N * K, K, 1]
if verbose:
print(
f"Changing gate_up_proj from {self.experts.gate_up_proj.size()}:{self.experts.gate_up_proj.stride()} to {permuted_shape}:{permuted_stride}"
)
with torch.no_grad():
self.experts.gate_up_proj.as_strided_(permuted_shape, permuted_stride)
if verbose:
print(
f"{self.experts.gate_up_proj.shape}:{self.experts.gate_up_proj.stride()}"
)
assert self.experts.down_proj.shape == torch.Size([E, N, K]), (
f"{self.experts.down_proj.shape} != {[E, N, K]}"
)
permuted_shape = [E, K, N]
permuted_stride = [K * N, N, 1]
if verbose:
print(
f"Changing down_proj from {self.experts.down_proj.size()}:{self.experts.down_proj.stride()} to {permuted_shape}:{permuted_stride}"
)
with torch.no_grad():
self.experts.down_proj.as_strided_(permuted_shape, permuted_stride)
if verbose:
print(f"{self.experts.down_proj.shape}:{self.experts.down_proj.stride()}")
if overlap_router_shared:
self.shared_expert_stream = torch.cuda.Stream()
self.default_event = torch.cuda.Event()
self.shared_expert_end_event = torch.cuda.Event()
@torch.no_grad
def copy_weights(self, other: Llama4TextMoe):
for name, param_to_copy in other.named_parameters():
if self.verbose:
print(f"Copying {name} with shape {param_to_copy.shape}")
param = self.get_parameter(name)
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
param_to_copy = param_to_copy.permute(0, 2, 1)
assert param.shape == param_to_copy.shape, (
f"{param.shape} != {param_to_copy.shape}"
)
param.copy_(param_to_copy)
return self
def check_weights(self, other: Llama4TextMoe):
for name, other_param in other.named_parameters():
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
other_param = other_param.permute(0, 2, 1)
param = self.get_parameter(name)
assert param.equal(other_param), f"Param {name} not equal!"
assert param.is_contiguous(), f"{name} not contiguous!"
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
assert x.shape[-1] == 2 * self.experts.expert_dim
gate_proj = x[..., : self.experts.expert_dim]
up_proj = x[..., self.experts.expert_dim :]
return self.experts.act_fn(gate_proj) * up_proj
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
# router_logits: (batch * sequence_length, n_experts)
hidden_states = hidden_states.view(-1, self.hidden_dim)
router_logits = self.router(hidden_states)
routing_weights, selected_experts = torch.topk(
router_logits, self.top_k, dim=-1
)
routing_weights = F.sigmoid(routing_weights.float()).to(hidden_states.dtype)
return router_logits, routing_weights, selected_experts
def get_token_counts_and_gather_indices(
self, selected_experts: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
token_counts_by_expert, gather_indices = get_routing_indices(
selected_experts, self.num_experts
)
assert not token_counts_by_expert.requires_grad
assert not gather_indices.requires_grad
return token_counts_by_expert, gather_indices
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
""" """
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k
hidden_states = hidden_states.view(-1, hidden_dim)
if self.overlap_router_shared:
# Marker for all prior ops on default stream
self.default_event.record()
router_logits, routing_weights, selected_experts = self.run_router(
hidden_states
)
assert routing_weights.shape == (num_tokens, self.top_k), (
f"{routing_weights.shape} != {(num_tokens, self.top_k)}"
)
if self.overlap_router_shared:
with torch.cuda.stream(self.shared_expert_stream):
# Ensure prior kernels on default stream complete
self.default_event.wait()
shared_expert_out = self.shared_expert(hidden_states)
# Ensure hidden states remains valid on this stream
hidden_states.record_stream(self.shared_expert_stream)
self.shared_expert_end_event.record()
# Ensure shared expert still valid on default stream
shared_expert_out.record_stream(torch.cuda.current_stream())
self.shared_expert_end_event.wait()
else:
shared_expert_out = self.shared_expert(hidden_states)
hidden_states = (
hidden_states.view(num_tokens, self.top_k, hidden_dim)
* routing_weights[..., None]
)
if self.top_k > 1:
hidden_states = hidden_states.sum(dim=1)
hidden_states_after_weight_merge = hidden_states.view(-1, hidden_dim)
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
token_counts_by_expert, gather_indices = (
self.get_token_counts_and_gather_indices(selected_experts)
)
# 2. Permute tokens from token order to expert order
hidden_states = permute(
hidden_states_after_weight_merge, gather_indices, self.top_k
)
assert hidden_states.shape == (total_tokens, hidden_dim)
# Start expert computation
first_gemm = torch_grouped_gemm(
X=hidden_states, W=self.experts.gate_up_proj, m_sizes=token_counts_by_expert
)
assert first_gemm.shape == (total_tokens, 2 * self.experts.expert_dim)
intermediate = self.act_and_mul(first_gemm)
assert intermediate.shape == (total_tokens, self.experts.expert_dim)
# See comment above
second_gemm = torch_grouped_gemm(
X=intermediate, W=self.experts.down_proj, m_sizes=token_counts_by_expert
)
assert second_gemm.shape == (total_tokens, hidden_dim)
# Post-processing
hidden_states_unpermute = unpermute(second_gemm, gather_indices)
assert hidden_states_unpermute.shape == (total_tokens, hidden_dim)
# grouped_gemm_out = hidden_states.view(batch_size, sequence_length, hidden_dim)
final_out = hidden_states_unpermute + shared_expert_out
result = (
Llama4MoeResult(
token_counts_by_expert=token_counts_by_expert,
gather_indices=gather_indices,
topk_weights=routing_weights,
hidden_states_after_weight_merge=hidden_states_after_weight_merge,
first_gemm=first_gemm,
intermediate=intermediate,
second_gemm=second_gemm,
hidden_states_unpermute=hidden_states_unpermute,
shared_expert_out=shared_expert_out,
final_out=final_out,
router_logits=router_logits,
)
if self.debug
else (final_out, routing_weights)
)
return result
class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
def __init__(
self,
config: Llama4TextConfig,
overlap_router_shared=False,
permute_x: bool = False,
permute_y: bool = True,
autotune: bool = True,
kernel_config_fwd: KernelConfigForward = None,
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
dW_only: bool = False,
dX_only: bool = False,
verbose=False,
):
super().__init__(config, overlap_router_shared=overlap_router_shared)
assert not permute_x, (
"Llama4 triton grouped gemm does not support permute x due to pre-multiplication of router weights"
)
self.permute_x = permute_x
self.permute_y = permute_y
self.autotune = autotune
if not autotune:
assert (
kernel_config_fwd is not None
and kernel_config_bwd_dW is not None
and kernel_config_bwd_dX is not None
), "Kernel configs must be provided if autotune is False"
self.kernel_config_fwd = kernel_config_fwd
self.kernel_config_bwd_dW = kernel_config_bwd_dW
self.kernel_config_bwd_dX = kernel_config_bwd_dX
self.dW_only = dW_only
self.dX_only = dX_only
@torch.no_grad
def copy_weights(self, other: Llama4TextMoe):
for name, param_to_copy in other.named_parameters():
if self.verbose:
print(f"Copying {name} with shape {param_to_copy.shape}")
param = self.get_parameter(name)
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
param_to_copy = param_to_copy.permute(0, 2, 1)
assert param.shape == param_to_copy.shape, (
f"{param.shape} != {param_to_copy.shape}"
)
param.copy_(param_to_copy)
return self
def check_weights(self, other: Llama4TextMoe):
for name, other_param in other.named_parameters():
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
other_param = other_param.permute(0, 2, 1)
param = self.get_parameter(name)
assert param.equal(other_param), f"Param {name} not equal!"
assert param.is_contiguous(), f"{name} not contiguous!"
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
assert x.shape[-1] == 2 * self.experts.expert_dim
gate_proj = x[..., : self.experts.expert_dim]
up_proj = x[..., self.experts.expert_dim :]
return self.experts.act_fn(gate_proj) * up_proj
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
# router_logits: (batch * sequence_length, n_experts)
hidden_states = hidden_states.view(-1, self.hidden_dim)
router_logits = self.router(hidden_states)
routing_weights, selected_experts = torch.topk(
router_logits, self.top_k, dim=-1
)
routing_weights = F.sigmoid(routing_weights.float()).to(hidden_states.dtype)
return router_logits, routing_weights, selected_experts
def get_token_counts_and_gather_indices(
self, selected_experts: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
token_counts_by_expert, gather_indices = get_routing_indices(
selected_experts, self.num_experts
)
assert not token_counts_by_expert.requires_grad
assert not gather_indices.requires_grad
return token_counts_by_expert, gather_indices
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
""" """
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k
hidden_states = hidden_states.view(-1, hidden_dim)
if self.overlap_router_shared:
# Marker for all prior ops on default stream
self.default_event.record()
router_logits, routing_weights, selected_experts = self.run_router(
hidden_states
)
assert routing_weights.shape == (num_tokens, self.top_k), (
f"{routing_weights.shape} != {(num_tokens, self.top_k)}"
)
if self.overlap_router_shared:
with torch.cuda.stream(self.shared_expert_stream):
# Ensure prior kernels on default stream complete
self.default_event.wait()
shared_expert_out = self.shared_expert(hidden_states)
# Ensure hidden states remains valid on this stream
hidden_states.record_stream(self.shared_expert_stream)
self.shared_expert_end_event.record()
# Ensure shared expert still valid on default stream
shared_expert_out.record_stream(torch.cuda.current_stream())
self.shared_expert_end_event.wait()
else:
shared_expert_out = self.shared_expert(hidden_states)
hidden_states = (
hidden_states.view(num_tokens, self.top_k, hidden_dim)
* routing_weights[..., None]
)
if self.top_k > 1:
hidden_states = hidden_states.sum(dim=1)
hidden_states = hidden_states.view(-1, hidden_dim)
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
token_counts_by_expert, gather_indices = (
self.get_token_counts_and_gather_indices(selected_experts)
)
# 2. Permute tokens from token order to expert order
hidden_states = permute(hidden_states, gather_indices, self.top_k)
assert hidden_states.shape == (total_tokens, hidden_dim)
# Start expert computation
hidden_states = grouped_gemm(
X=hidden_states,
W=self.experts.gate_up_proj,
m_sizes=token_counts_by_expert,
gather_indices=gather_indices,
topk=self.top_k,
permute_x=self.permute_x,
permute_y=False, # output of first grouped gemm should never be permuted
autotune=self.autotune,
kernel_config_fwd=self.kernel_config_fwd,
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
is_first_gemm=True,
dW_only=self.dW_only,
dX_only=self.dX_only,
)
hidden_states = self.act_and_mul(hidden_states)
hidden_states = grouped_gemm(
X=hidden_states,
W=self.experts.down_proj,
m_sizes=token_counts_by_expert,
gather_indices=gather_indices,
topk=self.top_k,
permute_x=False,
permute_y=self.permute_y,
autotune=self.autotune,
kernel_config_fwd=self.kernel_config_fwd,
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
is_first_gemm=False,
dW_only=self.dW_only,
dX_only=self.dX_only,
)
# Post-processing
# 1. Unpermute from expert order to token order
if not self.permute_y:
hidden_states = unpermute(hidden_states, gather_indices)
hidden_states += shared_expert_out
return hidden_states, routing_weights

View file

@ -0,0 +1,345 @@
from dataclasses import dataclass
from typing import Tuple
import torch
import torch.nn.functional as F
from transformers.models.qwen3_moe.configuration_qwen3_moe import Qwen3MoeConfig
from transformers.models.qwen3_moe.modeling_qwen3_moe import (
ACT2FN,
Qwen3MoeSparseMoeBlock,
)
from grouped_gemm.interface import grouped_gemm
from grouped_gemm.kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
get_routing_indices,
permute,
torch_grouped_gemm,
unpermute,
)
"""
Reference implementation of HF Qwen3 MoE block using grouped gemm.
The Qwen3MoeGroupedGEMMBlock is a reference torch-native implemention.
Qwen3MoeFusedGroupedGEMMBlock is a version using the triton grouped gemm kernel.
NOTE: This is NOT to be used for production as it contains many extra checks and saves all intermediate results for debugging.
"""
@dataclass
class GroupedGEMMResult:
token_counts_by_expert: torch.Tensor
gather_indices: torch.Tensor
topk_weights: torch.Tensor
first_gemm: torch.Tensor
intermediate: torch.Tensor
second_gemm: torch.Tensor
hidden_states_unpermute: torch.Tensor
hidden_states: torch.Tensor # final output
class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
def __init__(
self,
config,
gate: torch.Tensor,
gate_up_proj: torch.Tensor,
down_proj: torch.Tensor,
):
super().__init__()
self.num_experts = config.num_experts
self.top_k = config.num_experts_per_tok
self.norm_topk_prob = config.norm_topk_prob
self.hidden_size = config.hidden_size
self.moe_intermediate_size = config.moe_intermediate_size
assert gate.shape == (config.num_experts, config.hidden_size)
assert gate_up_proj.shape == (
config.num_experts,
2 * config.moe_intermediate_size,
config.hidden_size,
)
assert down_proj.shape == (
config.num_experts,
config.hidden_size,
config.moe_intermediate_size,
)
# gating
self.gate = torch.nn.Parameter(gate)
# experts
self.gate_up_proj = torch.nn.Parameter(gate_up_proj, requires_grad=True)
self.down_proj = torch.nn.Parameter(down_proj, requires_grad=True)
self.act_fn = ACT2FN[config.hidden_act]
@staticmethod
def extract_hf_weights(moe_block: Qwen3MoeSparseMoeBlock):
config: Qwen3MoeConfig = moe_block.experts[0].config
num_experts = config.num_experts
gate = moe_block.gate.weight.data
gate_proj = torch.stack(
[moe_block.experts[i].gate_proj.weight.data for i in range(num_experts)],
dim=0,
)
up_proj = torch.stack(
[moe_block.experts[i].up_proj.weight.data for i in range(num_experts)],
dim=0,
)
down_proj = torch.stack(
[moe_block.experts[i].down_proj.weight.data for i in range(num_experts)],
dim=0,
)
gate_up_proj = torch.cat([gate_proj, up_proj], dim=1)
return gate, gate_up_proj, down_proj
@classmethod
def from_hf(cls, moe_block: Qwen3MoeSparseMoeBlock):
config: Qwen3MoeConfig = moe_block.experts[0].config
gate, gate_up_proj, down_proj = cls.extract_hf_weights(moe_block)
return cls(config, gate, gate_up_proj, down_proj)
def check_weights(self, moe_block: Qwen3MoeSparseMoeBlock):
for i in range(self.num_experts):
assert self.gate_up_proj[i].equal(
torch.cat(
[
moe_block.experts[i].gate_proj.weight.data,
moe_block.experts[i].up_proj.weight.data,
],
dim=0,
)
)
assert self.down_proj[i].equal(moe_block.experts[i].down_proj.weight.data)
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
assert x.shape[-1] == 2 * self.moe_intermediate_size
gate_proj = x[..., : self.moe_intermediate_size]
up_proj = x[..., self.moe_intermediate_size :]
return self.act_fn(gate_proj) * up_proj
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
# router_logits: (batch * sequence_length, n_experts)
router_logits = torch.nn.functional.linear(hidden_states, self.gate)
routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float)
routing_weights, selected_experts = torch.topk(
routing_weights, self.top_k, dim=-1
)
if self.norm_topk_prob: # only diff with mixtral sparse moe block!
routing_weights /= routing_weights.sum(dim=-1, keepdim=True)
# we cast back to the input dtype
routing_weights = routing_weights.to(hidden_states.dtype)
return router_logits, routing_weights, selected_experts
def get_token_counts_and_gather_indices(
self, selected_experts: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
token_counts_by_expert, gather_indices = get_routing_indices(
selected_experts, self.num_experts
)
assert not token_counts_by_expert.requires_grad
assert not gather_indices.requires_grad
return token_counts_by_expert, gather_indices
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
""" """
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k
hidden_states = hidden_states.view(-1, hidden_dim)
router_logits, routing_weights, selected_experts = self.run_router(
hidden_states
)
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
token_counts_by_expert, gather_indices = (
self.get_token_counts_and_gather_indices(selected_experts)
)
# 2. Permute tokens from token order to expert order
hidden_states = permute(hidden_states, gather_indices, self.top_k)
assert hidden_states.shape == (total_tokens, hidden_dim)
# Start expert computation
first_gemm = torch_grouped_gemm(
X=hidden_states, W=self.gate_up_proj, m_sizes=token_counts_by_expert
)
assert first_gemm.shape == (total_tokens, 2 * self.moe_intermediate_size)
intermediate = self.act_and_mul(first_gemm)
assert intermediate.shape == (total_tokens, self.moe_intermediate_size)
second_gemm = torch_grouped_gemm(
X=intermediate, W=self.down_proj, m_sizes=token_counts_by_expert
)
assert second_gemm.shape == (total_tokens, hidden_dim)
# Post-processing
# 1. Unpermute from expert order to token order
hidden_states_unpermute = unpermute(second_gemm, gather_indices)
assert hidden_states_unpermute.shape == (total_tokens, hidden_dim)
# 2. Merge topk weights
hidden_states = (
hidden_states_unpermute.view(num_tokens, self.top_k, hidden_dim)
* routing_weights[..., None]
)
hidden_states = hidden_states.sum(dim=1)
assert hidden_states.shape == (num_tokens, hidden_dim)
hidden_states = hidden_states.view(batch_size, sequence_length, hidden_dim)
return GroupedGEMMResult(
token_counts_by_expert=token_counts_by_expert,
gather_indices=gather_indices,
topk_weights=routing_weights,
first_gemm=first_gemm,
intermediate=intermediate,
second_gemm=second_gemm,
hidden_states_unpermute=hidden_states_unpermute,
hidden_states=hidden_states,
), router_logits
class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock):
def __init__(
self,
config: Qwen3MoeConfig,
gate: torch.Tensor,
gate_up_proj: torch.Tensor,
down_proj: torch.Tensor,
permute_x: bool = True,
permute_y: bool = True,
autotune: bool = True,
kernel_config_fwd: KernelConfigForward = None,
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
dW_only: bool = False,
dX_only: bool = False,
):
super().__init__(config, gate, gate_up_proj, down_proj)
self.permute_x = permute_x
self.permute_y = permute_y
self.autotune = autotune
if not autotune:
assert (
kernel_config_fwd is not None
and kernel_config_bwd_dW is not None
and kernel_config_bwd_dX is not None
), "Kernel configs must be provided if autotune is False"
self.kernel_config_fwd = kernel_config_fwd
self.kernel_config_bwd_dW = kernel_config_bwd_dW
self.kernel_config_bwd_dX = kernel_config_bwd_dX
self.dW_only = dW_only
self.dX_only = dX_only
@classmethod
def from_hf(
cls,
moe_block: Qwen3MoeSparseMoeBlock,
permute_x: bool = True,
permute_y: bool = True,
autotune: bool = True,
kernel_config_fwd: KernelConfigForward = None,
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
dW_only: bool = False,
dX_only: bool = False,
):
config: Qwen3MoeConfig = moe_block.experts[0].config
gate, gate_up_proj, down_proj = Qwen3MoeGroupedGEMMBlock.extract_hf_weights(
moe_block
)
return cls(
config,
gate,
gate_up_proj,
down_proj,
permute_x=permute_x,
permute_y=permute_y,
autotune=autotune,
kernel_config_fwd=kernel_config_fwd,
kernel_config_bwd_dW=kernel_config_bwd_dW,
kernel_config_bwd_dX=kernel_config_bwd_dX,
dW_only=dW_only,
dX_only=dX_only,
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k
hidden_states = hidden_states.view(-1, hidden_dim)
router_logits, routing_weights, selected_experts = self.run_router(
hidden_states
)
# Pre-processing
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
token_counts_by_expert, gather_indices = (
self.get_token_counts_and_gather_indices(selected_experts)
)
# 2. permute_x -> permutation will be fused in prologue of first grouped gemm
if not self.permute_x:
hidden_states = permute(hidden_states, gather_indices, self.top_k)
# Start expert computation
hidden_states = grouped_gemm(
X=hidden_states,
W=self.gate_up_proj,
m_sizes=token_counts_by_expert,
gather_indices=gather_indices,
topk=self.top_k,
permute_x=self.permute_x,
permute_y=False, # output of first grouped gemm should never be permuted
autotune=self.autotune,
kernel_config_fwd=self.kernel_config_fwd,
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
is_first_gemm=True,
dW_only=self.dW_only,
dX_only=self.dX_only,
)
hidden_states = self.act_and_mul(hidden_states)
hidden_states = grouped_gemm(
X=hidden_states,
W=self.down_proj,
m_sizes=token_counts_by_expert,
gather_indices=gather_indices,
topk=self.top_k,
permute_x=False,
permute_y=self.permute_y,
autotune=self.autotune,
kernel_config_fwd=self.kernel_config_fwd,
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
is_first_gemm=False,
dW_only=self.dW_only,
dX_only=self.dX_only,
)
# Post-processing
# 1. Unpermute from expert order to token order
if not self.permute_y:
hidden_states = unpermute(hidden_states, gather_indices)
# 2. Merge topk weights
hidden_states = (
hidden_states.view(num_tokens, self.top_k, hidden_dim)
* routing_weights[..., None]
)
hidden_states = hidden_states.sum(dim=1)
hidden_states = hidden_states.view(batch_size, sequence_length, hidden_dim)
return hidden_states, router_logits

View file

@ -1,14 +1,6 @@
from dataclasses import dataclass
from typing import Tuple
import torch
import torch.nn as nn
import torch.nn.functional as F
from transformers.models.qwen3_moe import Qwen3MoeConfig
from transformers.models.qwen3_moe.modeling_qwen3_moe import (
ACT2FN,
Qwen3MoeSparseMoeBlock,
)
def permute(X: torch.Tensor, gather_indices: torch.Tensor, topk: int):
@ -58,13 +50,9 @@ def calculate_topk(
def _activation(gating_output: torch.Tensor):
if use_sigmoid:
scores = torch.sigmoid(gating_output.to(torch.float32)).to(
gating_output.dtype
)
scores = torch.sigmoid(gating_output.to(torch.float32)).to(gating_output.dtype)
else:
scores = F.softmax(gating_output.to(torch.float32), dim=1).to(
gating_output.dtype
)
scores = F.softmax(gating_output.to(torch.float32), dim=1).to(gating_output.dtype)
return scores
@ -79,17 +67,13 @@ def calculate_topk(
topk_weights = _activation(topk_weights)
if renormalize:
topk_weights /= torch.sum(topk_weights, dim=-1, keepdim=True).to(
gating_output.dtype
)
topk_weights /= torch.sum(topk_weights, dim=-1, keepdim=True).to(gating_output.dtype)
return topk_weights, topk_ids
@torch.no_grad()
def get_routing_indices(
selected_experts, num_experts, return_scatter_indices: bool = False
):
def get_routing_indices(selected_experts, num_experts, return_scatter_indices: bool = False):
"""
Returns:
token_counts_by_expert: [num_experts]
@ -155,181 +139,3 @@ def torch_grouped_gemm(X, W, m_sizes, transpose=True):
m_start = m_end
return result
@dataclass
class GroupedGEMMResult:
token_counts_by_expert: torch.Tensor
gather_indices: torch.Tensor
topk_weights: torch.Tensor
first_gemm: torch.Tensor
intermediate: torch.Tensor
second_gemm: torch.Tensor
hidden_states_unpermute: torch.Tensor
hidden_states: torch.Tensor # final output
class Qwen3MoeGroupedGEMMBlock(nn.Module):
def __init__(
self,
config,
gate: torch.Tensor,
gate_up_proj: torch.Tensor,
down_proj: torch.Tensor,
):
super().__init__()
self.num_experts = config.num_experts
self.top_k = config.num_experts_per_tok
self.norm_topk_prob = config.norm_topk_prob
self.hidden_size = config.hidden_size
self.moe_intermediate_size = config.moe_intermediate_size
assert gate.shape == (config.num_experts, config.hidden_size)
assert gate_up_proj.shape == (
config.num_experts,
2 * config.moe_intermediate_size,
config.hidden_size,
)
assert down_proj.shape == (
config.num_experts,
config.hidden_size,
config.moe_intermediate_size,
)
# gating
self.gate = torch.nn.Parameter(gate)
# experts
self.gate_up_proj = torch.nn.Parameter(gate_up_proj, requires_grad=True)
self.down_proj = torch.nn.Parameter(down_proj, requires_grad=True)
self.act_fn = ACT2FN[config.hidden_act]
@staticmethod
def extract_hf_weights(moe_block: Qwen3MoeSparseMoeBlock):
config: Qwen3MoeConfig = moe_block.experts[0].config
num_experts = config.num_experts
gate = moe_block.gate.weight.data
gate_proj = torch.stack(
[moe_block.experts[i].gate_proj.weight.data for i in range(num_experts)],
dim=0,
)
up_proj = torch.stack(
[moe_block.experts[i].up_proj.weight.data for i in range(num_experts)],
dim=0,
)
down_proj = torch.stack(
[moe_block.experts[i].down_proj.weight.data for i in range(num_experts)],
dim=0,
)
gate_up_proj = torch.cat([gate_proj, up_proj], dim=1)
return gate, gate_up_proj, down_proj
@classmethod
def from_hf(cls, moe_block: Qwen3MoeSparseMoeBlock):
config: Qwen3MoeConfig = moe_block.experts[0].config
gate, gate_up_proj, down_proj = cls.extract_hf_weights(moe_block)
return cls(config, gate, gate_up_proj, down_proj)
def check_weights(self, moe_block: Qwen3MoeSparseMoeBlock):
for i in range(self.num_experts):
assert self.gate_up_proj[i].equal(
torch.cat(
[
moe_block.experts[i].gate_proj.weight.data,
moe_block.experts[i].up_proj.weight.data,
],
dim=0,
)
)
assert self.down_proj[i].equal(moe_block.experts[i].down_proj.weight.data)
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
assert x.shape[-1] == 2 * self.moe_intermediate_size
gate_proj = x[..., : self.moe_intermediate_size]
up_proj = x[..., self.moe_intermediate_size :]
return self.act_fn(gate_proj) * up_proj
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
# router_logits: (batch * sequence_length, n_experts)
router_logits = torch.nn.functional.linear(hidden_states, self.gate)
routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float)
routing_weights, selected_experts = torch.topk(
routing_weights, self.top_k, dim=-1
)
if self.norm_topk_prob: # only diff with mixtral sparse moe block!
routing_weights /= routing_weights.sum(dim=-1, keepdim=True)
# we cast back to the input dtype
routing_weights = routing_weights.to(hidden_states.dtype)
return router_logits, routing_weights, selected_experts
def get_token_counts_and_gather_indices(
self, selected_experts: torch.Tensor
) -> Tuple[torch.Tensor, torch.Tensor]:
token_counts_by_expert, gather_indices = get_routing_indices(
selected_experts, self.num_experts
)
assert not token_counts_by_expert.requires_grad
assert not gather_indices.requires_grad
return token_counts_by_expert, gather_indices
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
""" """
batch_size, sequence_length, hidden_dim = hidden_states.shape
num_tokens = batch_size * sequence_length
total_tokens = num_tokens * self.top_k
hidden_states = hidden_states.view(-1, hidden_dim)
router_logits, routing_weights, selected_experts = self.run_router(
hidden_states
)
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
token_counts_by_expert, gather_indices = (
self.get_token_counts_and_gather_indices(selected_experts)
)
# 2. Permute tokens from token order to expert order
hidden_states = permute(hidden_states, gather_indices, self.top_k)
assert hidden_states.shape == (total_tokens, hidden_dim)
# Start expert computation
first_gemm = torch_grouped_gemm(
X=hidden_states, W=self.gate_up_proj, m_sizes=token_counts_by_expert
)
assert first_gemm.shape == (total_tokens, 2 * self.moe_intermediate_size)
intermediate = self.act_and_mul(first_gemm)
assert intermediate.shape == (total_tokens, self.moe_intermediate_size)
second_gemm = torch_grouped_gemm(
X=intermediate, W=self.down_proj, m_sizes=token_counts_by_expert
)
assert second_gemm.shape == (total_tokens, hidden_dim)
# Post-processing
# 1. Unpermute from expert order to token order
hidden_states_unpermute = unpermute(second_gemm, gather_indices)
assert hidden_states_unpermute.shape == (total_tokens, hidden_dim)
# 2. Merge topk weights
hidden_states = (
hidden_states_unpermute.view(num_tokens, self.top_k, hidden_dim)
* routing_weights[..., None]
)
hidden_states = hidden_states.sum(dim=1)
assert hidden_states.shape == (num_tokens, hidden_dim)
hidden_states = hidden_states.view(batch_size, sequence_length, hidden_dim)
return GroupedGEMMResult(
token_counts_by_expert=token_counts_by_expert,
gather_indices=gather_indices,
topk_weights=routing_weights,
first_gemm=first_gemm,
intermediate=intermediate,
second_gemm=second_gemm,
hidden_states_unpermute=hidden_states_unpermute,
hidden_states=hidden_states,
), router_logits

View file

@ -13,12 +13,11 @@ from grouped_gemm.kernels.tuning import (
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import (
from grouped_gemm.reference.layers.qwen3_moe import (
GroupedGEMMResult,
Qwen3MoeGroupedGEMMBlock,
permute,
unpermute,
)
from grouped_gemm.reference.moe_ops import permute, unpermute
def rebind_experts_to_shared_buffer(

View file

@ -0,0 +1,259 @@
import argparse
import sys
from contextlib import contextmanager
from functools import partial
import pytest
import torch
from transformers import AutoConfig
from transformers.models.llama4 import Llama4Config, Llama4TextConfig
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
from grouped_gemm.kernels.tuning import (
KernelConfigBackward_dW,
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.layers.llama4_moe import (
Llama4GroupedGemmTextMoe,
Llama4TritonTextMoe,
)
TOLERANCES = {
torch.bfloat16: (1e-2, 1e-2),
torch.float16: (1e-3, 1e-3),
torch.float: (1e-5, 1e-5),
}
LLAMA4_SCOUT_ID = "meta-llama/Llama-4-Scout-17B-16E"
SEED = 42
SEQ_LENS = [1024]
DTYPES = [torch.bfloat16]
# Reduce the number of autotuning configs to prevent excessive runtime
NUM_AUTOTUNE_CONFIGS = 50
@contextmanager
def annotated_context(prelude, epilogue="Passed!", char="-", num_chars=80):
print(char * num_chars)
print(prelude)
yield
print(epilogue)
print(char * num_chars)
def get_text_config(model_id):
config: Llama4Config = AutoConfig.from_pretrained(model_id)
return config.text_config
def prep_triton_kernel_traits(autotune):
if not autotune:
kernel_config_fwd = KernelConfigForward()
kernel_config_bwd_dW = KernelConfigBackward_dW()
kernel_config_bwd_dX = KernelConfigBackward_dX()
else:
from grouped_gemm.kernels.backward import (
_autotuned_grouped_gemm_dW_kernel,
_autotuned_grouped_gemm_dX_kernel,
)
from grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
# Hack to reduce number of autotuning configs
_autotuned_grouped_gemm_forward_kernel.configs = (
_autotuned_grouped_gemm_forward_kernel.configs[:NUM_AUTOTUNE_CONFIGS]
)
_autotuned_grouped_gemm_dW_kernel.configs = (
_autotuned_grouped_gemm_dW_kernel.configs[:NUM_AUTOTUNE_CONFIGS]
)
_autotuned_grouped_gemm_dX_kernel.configs = (
_autotuned_grouped_gemm_dX_kernel.configs[:NUM_AUTOTUNE_CONFIGS]
)
kernel_config_fwd = None
kernel_config_bwd_dW = None
kernel_config_bwd_dX = None
return kernel_config_fwd, kernel_config_bwd_dW, kernel_config_bwd_dX
def sparse_to_dense(t: torch.Tensor):
t = t.sum(dim=0).view(-1)
return t
@torch.no_grad()
def _check_diff(
t1: torch.Tensor,
t2: torch.Tensor,
atol,
rtol,
precision=".6f",
verbose=False,
msg="",
):
t2 = t2.view_as(t1)
diff = t1.sub(t2).abs().max().item()
if verbose:
if msg == "":
msg = "diff"
print(f"{msg}: {diff:{precision}}")
assert torch.allclose(t1, t2, atol=atol, rtol=rtol)
def run_backwards(y: torch.Tensor, grad_output: torch.Tensor, module: torch.nn.Module):
y.backward(grad_output)
for name, param in module.named_parameters():
assert param.grad is not None, f"{name} missing grad!"
def _check_grads(
m1: torch.nn.Module,
m2: torch.nn.Module,
atol,
rtol,
precision=".6f",
verbose=False,
msg="",
):
for name, param in m1.named_parameters():
_check_diff(
param.grad,
m2.get_parameter(name).grad,
atol=atol,
rtol=rtol,
precision=precision,
verbose=verbose,
msg=f"{msg}:{name}.grad",
)
@pytest.fixture
def model_config():
return AutoConfig.from_pretrained(LLAMA4_SCOUT_ID).text_config
@pytest.mark.parametrize(
"overlap_router_shared",
[False, True],
ids=lambda x: "overlap_router_shared" if x else "no_overlap",
)
@pytest.mark.parametrize(
"permute_y", [False, True], ids=lambda x: "permute_y" if x else "no_permute_y"
)
@pytest.mark.parametrize(
"permute_x", [False], ids=lambda x: "permute_x" if x else "no_permute_x"
) # Llama4 does not support permute_x
@pytest.mark.parametrize(
"autotune", [True], ids=lambda x: "autotune" if x else "manual"
)
@pytest.mark.parametrize("seqlen", SEQ_LENS, ids=lambda x: f"seqlen={x}")
@pytest.mark.parametrize("dtype", DTYPES, ids=str)
def test_llama4_ref(
dtype: torch.dtype,
seqlen,
autotune: bool,
permute_x: bool,
permute_y: bool,
overlap_router_shared: bool,
model_config: Llama4TextConfig, # test fixture
bs: int = 1,
device="cuda",
precision=".6f",
verbose=False,
):
torch.manual_seed(
SEED
) # Should not be needed when running using pytest -- autouse fixture in conftest.py
device = "cuda"
hidden_dim = model_config.hidden_size
atol, rtol = TOLERANCES[dtype]
check_diff = partial(
_check_diff, atol=atol, rtol=rtol, precision=precision, verbose=verbose
)
check_grads = partial(
_check_grads, atol=atol, rtol=rtol, precision=precision, verbose=verbose
)
# Reference op -- HF
llama4_ref = Llama4TextMoe(model_config).to(dtype=dtype, device=device)
# Torch grouped gemm impl
llama4_gg_ref = Llama4GroupedGemmTextMoe(
model_config, overlap_router_shared=overlap_router_shared
).to(dtype=dtype, device=device)
llama4_gg_ref.copy_weights(llama4_ref)
llama4_gg_ref.check_weights(llama4_ref)
x_ref = torch.randn(
bs, seqlen, hidden_dim, dtype=dtype, device=device, requires_grad=True
)
x_torch_gg = x_ref.detach().clone().requires_grad_()
x_triton = x_ref.detach().clone().requires_grad_()
y_ref, routing_ref = llama4_ref(x_ref)
y_torch_gg, routing_torch_gg = llama4_gg_ref(x_torch_gg)
assert y_ref.shape == y_torch_gg.shape, f"{y_ref.shape} != {y_torch_gg.shape}"
with annotated_context("Testing torch grouped gemm Llama4TextMoe"):
check_diff(y_ref, y_torch_gg, msg="y_torch_gg")
check_diff(
sparse_to_dense(routing_ref), routing_torch_gg, msg="routing_torch_gg"
)
kernel_config_fwd, kernel_config_bwd_dW, kernel_config_bwd_dX = (
prep_triton_kernel_traits(autotune)
)
llama4_triton = Llama4TritonTextMoe(
model_config,
overlap_router_shared=overlap_router_shared,
permute_x=permute_x,
permute_y=permute_y,
autotune=autotune,
kernel_config_fwd=kernel_config_fwd,
kernel_config_bwd_dW=kernel_config_bwd_dW,
kernel_config_bwd_dX=kernel_config_bwd_dX,
).to(device=device, dtype=dtype)
llama4_triton.copy_weights(llama4_ref)
llama4_triton.check_weights(llama4_ref)
y_triton, routing_triton = llama4_triton(x_triton)
with annotated_context("Testing triton grouped gemm Llama4TextMoe forward"):
check_diff(y_ref, y_triton, msg="y_triton")
check_diff(sparse_to_dense(routing_ref), routing_triton, msg="routing_triton")
ref_grad = torch.randn_like(y_ref)
run_backwards(y_ref, ref_grad, llama4_ref)
run_backwards(y_torch_gg, ref_grad, llama4_gg_ref)
with annotated_context("Testing torch group gemm Llama4TextMoe backward"):
check_grads(llama4_ref, llama4_gg_ref, msg="torch_gg")
run_backwards(y_triton, ref_grad, llama4_triton)
with annotated_context("Testing triton group gemm Llama4TextMoe backward"):
check_grads(llama4_ref, llama4_triton, msg="triton")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--seqlen", type=int, default=1024)
parser.add_argument(
"--dtype", type=str, choices=["bfloat16", "float16"], default="bfloat16"
)
args = parser.parse_args()
args.dtype = getattr(torch, args.dtype)
args_dict = vars(args)
model_id = LLAMA4_SCOUT_ID
text_config: Llama4TextConfig = get_text_config(model_id)
for overlap in [False, True]:
test_llama4_ref(
seqlen=args.seqlen,
model_config=text_config,
dtype=args.dtype,
autotune=True,
permute_x=False,
permute_y=True,
overlap_router_shared=overlap,
verbose=True,
)

View file

@ -12,7 +12,7 @@ from grouped_gemm.kernels.tuning import (
KernelConfigBackward_dX,
KernelConfigForward,
)
from grouped_gemm.reference.moe_ops import Qwen3MoeGroupedGEMMBlock
from grouped_gemm.reference.layers.qwen3_moe import Qwen3MoeGroupedGEMMBlock
from .moe_utils import (
Qwen3MoeFusedGroupedGEMMBlock,
@ -76,7 +76,7 @@ def config(model_id: str):
@contextmanager
def test_context(prelude, epilogue="Passed!", char="-", num_chars=80):
def annotated_context(prelude, epilogue="Passed!", char="-", num_chars=80):
print(char * num_chars)
print(prelude)
yield
@ -110,8 +110,6 @@ def test_qwen3_moe(
permute_x: bool,
permute_y: bool,
autotune: bool,
atol: float,
rtol: float,
):
torch.manual_seed(
SEED
@ -119,6 +117,7 @@ def test_qwen3_moe(
device = "cuda"
hidden_size = config.hidden_size
bs = 1
atol, rtol = TOLERANCES[dtype]
# Reference op -- HF
moe_block = Qwen3MoeSparseMoeBlock(config).to(device, dtype)
@ -173,7 +172,7 @@ def test_qwen3_moe(
grouped_result = run_forward(grouped_gemm_block, X, is_grouped_gemm=True)
fused_result = run_forward(fused_gemm_block, X, is_grouped_gemm=True)
with test_context(
with annotated_context(
"Testing forward pass",
epilogue="Passed forward tests!",
char="=",
@ -181,10 +180,12 @@ def test_qwen3_moe(
):
# Sanity checks
with test_context("Checking HF vs torch grouped gemm MoE forward outputs..."):
with annotated_context(
"Checking HF vs torch grouped gemm MoE forward outputs..."
):
check_fwd(ref_result, grouped_result, atol, rtol, verbose=False)
with test_context(
with annotated_context(
"Checking torch grouped gemm MoE vs fused grouped gemm MoE forward outputs..."
):
# We implement a custom check for grouped gemm results to test each of the intermediate results for easier debugging
@ -197,7 +198,9 @@ def test_qwen3_moe(
verbose=False,
)
# Actual test
with test_context("Checking HF vs fused grouped gemm MoE forward outputs..."):
with annotated_context(
"Checking HF vs fused grouped gemm MoE forward outputs..."
):
check_fwd(ref_result, fused_result, atol, rtol, verbose=True)
# Backward
@ -215,18 +218,18 @@ def test_qwen3_moe(
fused_gemm_block, grad_output, output=fused_result.output, X=fused_result.X
)
with test_context(
with annotated_context(
"Testing backward pass",
epilogue="Passed backward tests!",
char="=",
num_chars=100,
):
# Sanity checks
with test_context("Checking HF vs torch grouped gemm MoE grads..."):
with annotated_context("Checking HF vs torch grouped gemm MoE grads..."):
check_grads(
ref_backward_result, grouped_backward_result, atol, rtol, verbose=False
)
with test_context(
with annotated_context(
"Checking torch grouped gemm MoE vs fused grouped gemm MoE grads..."
):
check_grads(
@ -238,7 +241,7 @@ def test_qwen3_moe(
)
# Actual test
with test_context("Checking HF vs fused grouped gemm MoE grads..."):
with annotated_context("Checking HF vs fused grouped gemm MoE grads..."):
check_grads(
ref_backward_result, fused_backward_result, atol, rtol, verbose=True
)
@ -264,4 +267,4 @@ if __name__ == "__main__":
print(
f"Testing {model_id} with seqlen={args.seqlen}, dtype={args.dtype}, permute_x={args.permute_x}, permute_y={args.permute_y}, autotune={args.autotune}, atol={atol}, rtol={rtol}"
)
test_qwen3_moe(config, atol=atol, rtol=rtol, **args_dict)
test_qwen3_moe(config, **args_dict)