diff --git a/unsloth/kernels/moe/README.md b/unsloth/kernels/moe/README.md
index 5cf342b4ca..326ee159c9 100644
--- a/unsloth/kernels/moe/README.md
+++ b/unsloth/kernels/moe/README.md
@@ -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
\ No newline at end of file
+ - Threadblock swizzling for better L2 caching
+ - Llama4
+ - Fused gather / topk weight merging
+ - Custom topk, gather indices kernel
+ - Shared expert fusion with experts calculation
\ No newline at end of file
diff --git a/unsloth/kernels/moe/benchmark/benchmark_fused_moe.py b/unsloth/kernels/moe/benchmark/benchmark_fused_moe.py
index 9d95e67541..2fe2afa1e8 100644
--- a/unsloth/kernels/moe/benchmark/benchmark_fused_moe.py
+++ b/unsloth/kernels/moe/benchmark/benchmark_fused_moe.py
@@ -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
+ )
diff --git a/unsloth/kernels/moe/benchmark/utils.py b/unsloth/kernels/moe/benchmark/utils.py
index 04c7e8aba3..3c4fd705df 100644
--- a/unsloth/kernels/moe/benchmark/utils.py
+++ b/unsloth/kernels/moe/benchmark/utils.py
@@ -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,
+ )
diff --git a/unsloth/kernels/moe/grouped_gemm/LICENSE b/unsloth/kernels/moe/grouped_gemm/LICENSE
new file mode 100644
index 0000000000..29ebfa545f
--- /dev/null
+++ b/unsloth/kernels/moe/grouped_gemm/LICENSE
@@ -0,0 +1,661 @@
+ GNU AFFERO GENERAL PUBLIC LICENSE
+ Version 3, 19 November 2007
+
+ Copyright (C) 2007 Free Software Foundation, Inc.
+ Everyone is permitted to copy and distribute verbatim copies
+ of this license document, but changing it is not allowed.
+
+ Preamble
+
+ The GNU Affero General Public License is a free, copyleft license for
+software and other kinds of works, specifically designed to ensure
+cooperation with the community in the case of network server software.
+
+ The licenses for most software and other practical works are designed
+to take away your freedom to share and change the works. By contrast,
+our General Public Licenses are intended to guarantee your freedom to
+share and change all versions of a program--to make sure it remains free
+software for all its users.
+
+ When we speak of free software, we are referring to freedom, not
+price. Our General Public Licenses are designed to make sure that you
+have the freedom to distribute copies of free software (and charge for
+them if you wish), that you receive source code or can get it if you
+want it, that you can change the software or use pieces of it in new
+free programs, and that you know you can do these things.
+
+ Developers that use our General Public Licenses protect your rights
+with two steps: (1) assert copyright on the software, and (2) offer
+you this License which gives you legal permission to copy, distribute
+and/or modify the software.
+
+ A secondary benefit of defending all users' freedom is that
+improvements made in alternate versions of the program, if they
+receive widespread use, become available for other developers to
+incorporate. Many developers of free software are heartened and
+encouraged by the resulting cooperation. However, in the case of
+software used on network servers, this result may fail to come about.
+The GNU General Public License permits making a modified version and
+letting the public access it on a server without ever releasing its
+source code to the public.
+
+ The GNU Affero General Public License is designed specifically to
+ensure that, in such cases, the modified source code becomes available
+to the community. It requires the operator of a network server to
+provide the source code of the modified version running there to the
+users of that server. Therefore, public use of a modified version, on
+a publicly accessible server, gives the public access to the source
+code of the modified version.
+
+ An older license, called the Affero General Public License and
+published by Affero, was designed to accomplish similar goals. This is
+a different license, not a version of the Affero GPL, but Affero has
+released a new version of the Affero GPL which permits relicensing under
+this license.
+
+ The precise terms and conditions for copying, distribution and
+modification follow.
+
+ TERMS AND CONDITIONS
+
+ 0. Definitions.
+
+ "This License" refers to version 3 of the GNU Affero General Public License.
+
+ "Copyright" also means copyright-like laws that apply to other kinds of
+works, such as semiconductor masks.
+
+ "The Program" refers to any copyrightable work licensed under this
+License. Each licensee is addressed as "you". "Licensees" and
+"recipients" may be individuals or organizations.
+
+ To "modify" a work means to copy from or adapt all or part of the work
+in a fashion requiring copyright permission, other than the making of an
+exact copy. The resulting work is called a "modified version" of the
+earlier work or a work "based on" the earlier work.
+
+ A "covered work" means either the unmodified Program or a work based
+on the Program.
+
+ To "propagate" a work means to do anything with it that, without
+permission, would make you directly or secondarily liable for
+infringement under applicable copyright law, except executing it on a
+computer or modifying a private copy. Propagation includes copying,
+distribution (with or without modification), making available to the
+public, and in some countries other activities as well.
+
+ To "convey" a work means any kind of propagation that enables other
+parties to make or receive copies. Mere interaction with a user through
+a computer network, with no transfer of a copy, is not conveying.
+
+ An interactive user interface displays "Appropriate Legal Notices"
+to the extent that it includes a convenient and prominently visible
+feature that (1) displays an appropriate copyright notice, and (2)
+tells the user that there is no warranty for the work (except to the
+extent that warranties are provided), that licensees may convey the
+work under this License, and how to view a copy of this License. If
+the interface presents a list of user commands or options, such as a
+menu, a prominent item in the list meets this criterion.
+
+ 1. Source Code.
+
+ The "source code" for a work means the preferred form of the work
+for making modifications to it. "Object code" means any non-source
+form of a work.
+
+ A "Standard Interface" means an interface that either is an official
+standard defined by a recognized standards body, or, in the case of
+interfaces specified for a particular programming language, one that
+is widely used among developers working in that language.
+
+ The "System Libraries" of an executable work include anything, other
+than the work as a whole, that (a) is included in the normal form of
+packaging a Major Component, but which is not part of that Major
+Component, and (b) serves only to enable use of the work with that
+Major Component, or to implement a Standard Interface for which an
+implementation is available to the public in source code form. A
+"Major Component", in this context, means a major essential component
+(kernel, window system, and so on) of the specific operating system
+(if any) on which the executable work runs, or a compiler used to
+produce the work, or an object code interpreter used to run it.
+
+ The "Corresponding Source" for a work in object code form means all
+the source code needed to generate, install, and (for an executable
+work) run the object code and to modify the work, including scripts to
+control those activities. However, it does not include the work's
+System Libraries, or general-purpose tools or generally available free
+programs which are used unmodified in performing those activities but
+which are not part of the work. For example, Corresponding Source
+includes interface definition files associated with source files for
+the work, and the source code for shared libraries and dynamically
+linked subprograms that the work is specifically designed to require,
+such as by intimate data communication or control flow between those
+subprograms and other parts of the work.
+
+ The Corresponding Source need not include anything that users
+can regenerate automatically from other parts of the Corresponding
+Source.
+
+ The Corresponding Source for a work in source code form is that
+same work.
+
+ 2. Basic Permissions.
+
+ All rights granted under this License are granted for the term of
+copyright on the Program, and are irrevocable provided the stated
+conditions are met. This License explicitly affirms your unlimited
+permission to run the unmodified Program. The output from running a
+covered work is covered by this License only if the output, given its
+content, constitutes a covered work. This License acknowledges your
+rights of fair use or other equivalent, as provided by copyright law.
+
+ You may make, run and propagate covered works that you do not
+convey, without conditions so long as your license otherwise remains
+in force. You may convey covered works to others for the sole purpose
+of having them make modifications exclusively for you, or provide you
+with facilities for running those works, provided that you comply with
+the terms of this License in conveying all material for which you do
+not control copyright. Those thus making or running the covered works
+for you must do so exclusively on your behalf, under your direction
+and control, on terms that prohibit them from making any copies of
+your copyrighted material outside their relationship with you.
+
+ Conveying under any other circumstances is permitted solely under
+the conditions stated below. Sublicensing is not allowed; section 10
+makes it unnecessary.
+
+ 3. Protecting Users' Legal Rights From Anti-Circumvention Law.
+
+ No covered work shall be deemed part of an effective technological
+measure under any applicable law fulfilling obligations under article
+11 of the WIPO copyright treaty adopted on 20 December 1996, or
+similar laws prohibiting or restricting circumvention of such
+measures.
+
+ When you convey a covered work, you waive any legal power to forbid
+circumvention of technological measures to the extent such circumvention
+is effected by exercising rights under this License with respect to
+the covered work, and you disclaim any intention to limit operation or
+modification of the work as a means of enforcing, against the work's
+users, your or third parties' legal rights to forbid circumvention of
+technological measures.
+
+ 4. Conveying Verbatim Copies.
+
+ You may convey verbatim copies of the Program's source code as you
+receive it, in any medium, provided that you conspicuously and
+appropriately publish on each copy an appropriate copyright notice;
+keep intact all notices stating that this License and any
+non-permissive terms added in accord with section 7 apply to the code;
+keep intact all notices of the absence of any warranty; and give all
+recipients a copy of this License along with the Program.
+
+ You may charge any price or no price for each copy that you convey,
+and you may offer support or warranty protection for a fee.
+
+ 5. Conveying Modified Source Versions.
+
+ You may convey a work based on the Program, or the modifications to
+produce it from the Program, in the form of source code under the
+terms of section 4, provided that you also meet all of these conditions:
+
+ a) The work must carry prominent notices stating that you modified
+ it, and giving a relevant date.
+
+ b) The work must carry prominent notices stating that it is
+ released under this License and any conditions added under section
+ 7. This requirement modifies the requirement in section 4 to
+ "keep intact all notices".
+
+ c) You must license the entire work, as a whole, under this
+ License to anyone who comes into possession of a copy. This
+ License will therefore apply, along with any applicable section 7
+ additional terms, to the whole of the work, and all its parts,
+ regardless of how they are packaged. This License gives no
+ permission to license the work in any other way, but it does not
+ invalidate such permission if you have separately received it.
+
+ d) If the work has interactive user interfaces, each must display
+ Appropriate Legal Notices; however, if the Program has interactive
+ interfaces that do not display Appropriate Legal Notices, your
+ work need not make them do so.
+
+ A compilation of a covered work with other separate and independent
+works, which are not by their nature extensions of the covered work,
+and which are not combined with it such as to form a larger program,
+in or on a volume of a storage or distribution medium, is called an
+"aggregate" if the compilation and its resulting copyright are not
+used to limit the access or legal rights of the compilation's users
+beyond what the individual works permit. Inclusion of a covered work
+in an aggregate does not cause this License to apply to the other
+parts of the aggregate.
+
+ 6. Conveying Non-Source Forms.
+
+ You may convey a covered work in object code form under the terms
+of sections 4 and 5, provided that you also convey the
+machine-readable Corresponding Source under the terms of this License,
+in one of these ways:
+
+ a) Convey the object code in, or embodied in, a physical product
+ (including a physical distribution medium), accompanied by the
+ Corresponding Source fixed on a durable physical medium
+ customarily used for software interchange.
+
+ b) Convey the object code in, or embodied in, a physical product
+ (including a physical distribution medium), accompanied by a
+ written offer, valid for at least three years and valid for as
+ long as you offer spare parts or customer support for that product
+ model, to give anyone who possesses the object code either (1) a
+ copy of the Corresponding Source for all the software in the
+ product that is covered by this License, on a durable physical
+ medium customarily used for software interchange, for a price no
+ more than your reasonable cost of physically performing this
+ conveying of source, or (2) access to copy the
+ Corresponding Source from a network server at no charge.
+
+ c) Convey individual copies of the object code with a copy of the
+ written offer to provide the Corresponding Source. This
+ alternative is allowed only occasionally and noncommercially, and
+ only if you received the object code with such an offer, in accord
+ with subsection 6b.
+
+ d) Convey the object code by offering access from a designated
+ place (gratis or for a charge), and offer equivalent access to the
+ Corresponding Source in the same way through the same place at no
+ further charge. You need not require recipients to copy the
+ Corresponding Source along with the object code. If the place to
+ copy the object code is a network server, the Corresponding Source
+ may be on a different server (operated by you or a third party)
+ that supports equivalent copying facilities, provided you maintain
+ clear directions next to the object code saying where to find the
+ Corresponding Source. Regardless of what server hosts the
+ Corresponding Source, you remain obligated to ensure that it is
+ available for as long as needed to satisfy these requirements.
+
+ e) Convey the object code using peer-to-peer transmission, provided
+ you inform other peers where the object code and Corresponding
+ Source of the work are being offered to the general public at no
+ charge under subsection 6d.
+
+ A separable portion of the object code, whose source code is excluded
+from the Corresponding Source as a System Library, need not be
+included in conveying the object code work.
+
+ A "User Product" is either (1) a "consumer product", which means any
+tangible personal property which is normally used for personal, family,
+or household purposes, or (2) anything designed or sold for incorporation
+into a dwelling. In determining whether a product is a consumer product,
+doubtful cases shall be resolved in favor of coverage. For a particular
+product received by a particular user, "normally used" refers to a
+typical or common use of that class of product, regardless of the status
+of the particular user or of the way in which the particular user
+actually uses, or expects or is expected to use, the product. A product
+is a consumer product regardless of whether the product has substantial
+commercial, industrial or non-consumer uses, unless such uses represent
+the only significant mode of use of the product.
+
+ "Installation Information" for a User Product means any methods,
+procedures, authorization keys, or other information required to install
+and execute modified versions of a covered work in that User Product from
+a modified version of its Corresponding Source. The information must
+suffice to ensure that the continued functioning of the modified object
+code is in no case prevented or interfered with solely because
+modification has been made.
+
+ If you convey an object code work under this section in, or with, or
+specifically for use in, a User Product, and the conveying occurs as
+part of a transaction in which the right of possession and use of the
+User Product is transferred to the recipient in perpetuity or for a
+fixed term (regardless of how the transaction is characterized), the
+Corresponding Source conveyed under this section must be accompanied
+by the Installation Information. But this requirement does not apply
+if neither you nor any third party retains the ability to install
+modified object code on the User Product (for example, the work has
+been installed in ROM).
+
+ The requirement to provide Installation Information does not include a
+requirement to continue to provide support service, warranty, or updates
+for a work that has been modified or installed by the recipient, or for
+the User Product in which it has been modified or installed. Access to a
+network may be denied when the modification itself materially and
+adversely affects the operation of the network or violates the rules and
+protocols for communication across the network.
+
+ Corresponding Source conveyed, and Installation Information provided,
+in accord with this section must be in a format that is publicly
+documented (and with an implementation available to the public in
+source code form), and must require no special password or key for
+unpacking, reading or copying.
+
+ 7. Additional Terms.
+
+ "Additional permissions" are terms that supplement the terms of this
+License by making exceptions from one or more of its conditions.
+Additional permissions that are applicable to the entire Program shall
+be treated as though they were included in this License, to the extent
+that they are valid under applicable law. If additional permissions
+apply only to part of the Program, that part may be used separately
+under those permissions, but the entire Program remains governed by
+this License without regard to the additional permissions.
+
+ When you convey a copy of a covered work, you may at your option
+remove any additional permissions from that copy, or from any part of
+it. (Additional permissions may be written to require their own
+removal in certain cases when you modify the work.) You may place
+additional permissions on material, added by you to a covered work,
+for which you have or can give appropriate copyright permission.
+
+ Notwithstanding any other provision of this License, for material you
+add to a covered work, you may (if authorized by the copyright holders of
+that material) supplement the terms of this License with terms:
+
+ a) Disclaiming warranty or limiting liability differently from the
+ terms of sections 15 and 16 of this License; or
+
+ b) Requiring preservation of specified reasonable legal notices or
+ author attributions in that material or in the Appropriate Legal
+ Notices displayed by works containing it; or
+
+ c) Prohibiting misrepresentation of the origin of that material, or
+ requiring that modified versions of such material be marked in
+ reasonable ways as different from the original version; or
+
+ d) Limiting the use for publicity purposes of names of licensors or
+ authors of the material; or
+
+ e) Declining to grant rights under trademark law for use of some
+ trade names, trademarks, or service marks; or
+
+ f) Requiring indemnification of licensors and authors of that
+ material by anyone who conveys the material (or modified versions of
+ it) with contractual assumptions of liability to the recipient, for
+ any liability that these contractual assumptions directly impose on
+ those licensors and authors.
+
+ All other non-permissive additional terms are considered "further
+restrictions" within the meaning of section 10. If the Program as you
+received it, or any part of it, contains a notice stating that it is
+governed by this License along with a term that is a further
+restriction, you may remove that term. If a license document contains
+a further restriction but permits relicensing or conveying under this
+License, you may add to a covered work material governed by the terms
+of that license document, provided that the further restriction does
+not survive such relicensing or conveying.
+
+ If you add terms to a covered work in accord with this section, you
+must place, in the relevant source files, a statement of the
+additional terms that apply to those files, or a notice indicating
+where to find the applicable terms.
+
+ Additional terms, permissive or non-permissive, may be stated in the
+form of a separately written license, or stated as exceptions;
+the above requirements apply either way.
+
+ 8. Termination.
+
+ You may not propagate or modify a covered work except as expressly
+provided under this License. Any attempt otherwise to propagate or
+modify it is void, and will automatically terminate your rights under
+this License (including any patent licenses granted under the third
+paragraph of section 11).
+
+ However, if you cease all violation of this License, then your
+license from a particular copyright holder is reinstated (a)
+provisionally, unless and until the copyright holder explicitly and
+finally terminates your license, and (b) permanently, if the copyright
+holder fails to notify you of the violation by some reasonable means
+prior to 60 days after the cessation.
+
+ Moreover, your license from a particular copyright holder is
+reinstated permanently if the copyright holder notifies you of the
+violation by some reasonable means, this is the first time you have
+received notice of violation of this License (for any work) from that
+copyright holder, and you cure the violation prior to 30 days after
+your receipt of the notice.
+
+ Termination of your rights under this section does not terminate the
+licenses of parties who have received copies or rights from you under
+this License. If your rights have been terminated and not permanently
+reinstated, you do not qualify to receive new licenses for the same
+material under section 10.
+
+ 9. Acceptance Not Required for Having Copies.
+
+ You are not required to accept this License in order to receive or
+run a copy of the Program. Ancillary propagation of a covered work
+occurring solely as a consequence of using peer-to-peer transmission
+to receive a copy likewise does not require acceptance. However,
+nothing other than this License grants you permission to propagate or
+modify any covered work. These actions infringe copyright if you do
+not accept this License. Therefore, by modifying or propagating a
+covered work, you indicate your acceptance of this License to do so.
+
+ 10. Automatic Licensing of Downstream Recipients.
+
+ Each time you convey a covered work, the recipient automatically
+receives a license from the original licensors, to run, modify and
+propagate that work, subject to this License. You are not responsible
+for enforcing compliance by third parties with this License.
+
+ An "entity transaction" is a transaction transferring control of an
+organization, or substantially all assets of one, or subdividing an
+organization, or merging organizations. If propagation of a covered
+work results from an entity transaction, each party to that
+transaction who receives a copy of the work also receives whatever
+licenses to the work the party's predecessor in interest had or could
+give under the previous paragraph, plus a right to possession of the
+Corresponding Source of the work from the predecessor in interest, if
+the predecessor has it or can get it with reasonable efforts.
+
+ You may not impose any further restrictions on the exercise of the
+rights granted or affirmed under this License. For example, you may
+not impose a license fee, royalty, or other charge for exercise of
+rights granted under this License, and you may not initiate litigation
+(including a cross-claim or counterclaim in a lawsuit) alleging that
+any patent claim is infringed by making, using, selling, offering for
+sale, or importing the Program or any portion of it.
+
+ 11. Patents.
+
+ A "contributor" is a copyright holder who authorizes use under this
+License of the Program or a work on which the Program is based. The
+work thus licensed is called the contributor's "contributor version".
+
+ A contributor's "essential patent claims" are all patent claims
+owned or controlled by the contributor, whether already acquired or
+hereafter acquired, that would be infringed by some manner, permitted
+by this License, of making, using, or selling its contributor version,
+but do not include claims that would be infringed only as a
+consequence of further modification of the contributor version. For
+purposes of this definition, "control" includes the right to grant
+patent sublicenses in a manner consistent with the requirements of
+this License.
+
+ Each contributor grants you a non-exclusive, worldwide, royalty-free
+patent license under the contributor's essential patent claims, to
+make, use, sell, offer for sale, import and otherwise run, modify and
+propagate the contents of its contributor version.
+
+ In the following three paragraphs, a "patent license" is any express
+agreement or commitment, however denominated, not to enforce a patent
+(such as an express permission to practice a patent or covenant not to
+sue for patent infringement). To "grant" such a patent license to a
+party means to make such an agreement or commitment not to enforce a
+patent against the party.
+
+ If you convey a covered work, knowingly relying on a patent license,
+and the Corresponding Source of the work is not available for anyone
+to copy, free of charge and under the terms of this License, through a
+publicly available network server or other readily accessible means,
+then you must either (1) cause the Corresponding Source to be so
+available, or (2) arrange to deprive yourself of the benefit of the
+patent license for this particular work, or (3) arrange, in a manner
+consistent with the requirements of this License, to extend the patent
+license to downstream recipients. "Knowingly relying" means you have
+actual knowledge that, but for the patent license, your conveying the
+covered work in a country, or your recipient's use of the covered work
+in a country, would infringe one or more identifiable patents in that
+country that you have reason to believe are valid.
+
+ If, pursuant to or in connection with a single transaction or
+arrangement, you convey, or propagate by procuring conveyance of, a
+covered work, and grant a patent license to some of the parties
+receiving the covered work authorizing them to use, propagate, modify
+or convey a specific copy of the covered work, then the patent license
+you grant is automatically extended to all recipients of the covered
+work and works based on it.
+
+ A patent license is "discriminatory" if it does not include within
+the scope of its coverage, prohibits the exercise of, or is
+conditioned on the non-exercise of one or more of the rights that are
+specifically granted under this License. You may not convey a covered
+work if you are a party to an arrangement with a third party that is
+in the business of distributing software, under which you make payment
+to the third party based on the extent of your activity of conveying
+the work, and under which the third party grants, to any of the
+parties who would receive the covered work from you, a discriminatory
+patent license (a) in connection with copies of the covered work
+conveyed by you (or copies made from those copies), or (b) primarily
+for and in connection with specific products or compilations that
+contain the covered work, unless you entered into that arrangement,
+or that patent license was granted, prior to 28 March 2007.
+
+ Nothing in this License shall be construed as excluding or limiting
+any implied license or other defenses to infringement that may
+otherwise be available to you under applicable patent law.
+
+ 12. No Surrender of Others' Freedom.
+
+ If conditions are imposed on you (whether by court order, agreement or
+otherwise) that contradict the conditions of this License, they do not
+excuse you from the conditions of this License. If you cannot convey a
+covered work so as to satisfy simultaneously your obligations under this
+License and any other pertinent obligations, then as a consequence you may
+not convey it at all. For example, if you agree to terms that obligate you
+to collect a royalty for further conveying from those to whom you convey
+the Program, the only way you could satisfy both those terms and this
+License would be to refrain entirely from conveying the Program.
+
+ 13. Remote Network Interaction; Use with the GNU General Public License.
+
+ Notwithstanding any other provision of this License, if you modify the
+Program, your modified version must prominently offer all users
+interacting with it remotely through a computer network (if your version
+supports such interaction) an opportunity to receive the Corresponding
+Source of your version by providing access to the Corresponding Source
+from a network server at no charge, through some standard or customary
+means of facilitating copying of software. This Corresponding Source
+shall include the Corresponding Source for any work covered by version 3
+of the GNU General Public License that is incorporated pursuant to the
+following paragraph.
+
+ Notwithstanding any other provision of this License, you have
+permission to link or combine any covered work with a work licensed
+under version 3 of the GNU General Public License into a single
+combined work, and to convey the resulting work. The terms of this
+License will continue to apply to the part which is the covered work,
+but the work with which it is combined will remain governed by version
+3 of the GNU General Public License.
+
+ 14. Revised Versions of this License.
+
+ The Free Software Foundation may publish revised and/or new versions of
+the GNU Affero General Public License from time to time. Such new versions
+will be similar in spirit to the present version, but may differ in detail to
+address new problems or concerns.
+
+ Each version is given a distinguishing version number. If the
+Program specifies that a certain numbered version of the GNU Affero General
+Public License "or any later version" applies to it, you have the
+option of following the terms and conditions either of that numbered
+version or of any later version published by the Free Software
+Foundation. If the Program does not specify a version number of the
+GNU Affero General Public License, you may choose any version ever published
+by the Free Software Foundation.
+
+ If the Program specifies that a proxy can decide which future
+versions of the GNU Affero General Public License can be used, that proxy's
+public statement of acceptance of a version permanently authorizes you
+to choose that version for the Program.
+
+ Later license versions may give you additional or different
+permissions. However, no additional obligations are imposed on any
+author or copyright holder as a result of your choosing to follow a
+later version.
+
+ 15. Disclaimer of Warranty.
+
+ THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY
+APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT
+HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY
+OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO,
+THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR
+PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM
+IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF
+ALL NECESSARY SERVICING, REPAIR OR CORRECTION.
+
+ 16. Limitation of Liability.
+
+ IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING
+WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS
+THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY
+GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE
+USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF
+DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD
+PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS),
+EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF
+SUCH DAMAGES.
+
+ 17. Interpretation of Sections 15 and 16.
+
+ If the disclaimer of warranty and limitation of liability provided
+above cannot be given local legal effect according to their terms,
+reviewing courts shall apply local law that most closely approximates
+an absolute waiver of all civil liability in connection with the
+Program, unless a warranty or assumption of liability accompanies a
+copy of the Program in return for a fee.
+
+ END OF TERMS AND CONDITIONS
+
+ How to Apply These Terms to Your New Programs
+
+ If you develop a new program, and you want it to be of the greatest
+possible use to the public, the best way to achieve this is to make it
+free software which everyone can redistribute and change under these terms.
+
+ To do so, attach the following notices to the program. It is safest
+to attach them to the start of each source file to most effectively
+state the exclusion of warranty; and each file should have at least
+the "copyright" line and a pointer to where the full notice is found.
+
+
+ Copyright (C)
+
+ This program is free software: you can redistribute it and/or modify
+ it under the terms of the GNU Affero General Public License as published
+ by the Free Software Foundation, either version 3 of the License, or
+ (at your option) any later version.
+
+ This program is distributed in the hope that it will be useful,
+ but WITHOUT ANY WARRANTY; without even the implied warranty of
+ MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
+ GNU Affero General Public License for more details.
+
+ You should have received a copy of the GNU Affero General Public License
+ along with this program. If not, see .
+
+Also add information on how to contact you by electronic and paper mail.
+
+ If your software can interact with users remotely through a computer
+network, you should also make sure that it provides a way for users to
+get its source. For example, if your program is a web application, its
+interface could display a "Source" link that leads users to an archive
+of the code. There are many ways you could offer source, and different
+solutions will be better for different programs; see section 13 for the
+specific requirements.
+
+ You should also get your employer (if you work as a programmer) or school,
+if any, to sign a "copyright disclaimer" for the program, if necessary.
+For more information on this, and how to apply and follow the GNU AGPL, see
+.
\ No newline at end of file
diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py
new file mode 100644
index 0000000000..63e64af858
--- /dev/null
+++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py
@@ -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
diff --git a/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py
new file mode 100644
index 0000000000..2bc9cc624d
--- /dev/null
+++ b/unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py
@@ -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
diff --git a/unsloth/kernels/moe/grouped_gemm/reference/moe_ops.py b/unsloth/kernels/moe/grouped_gemm/reference/moe_ops.py
index 4dd97d0958..08c04778a0 100644
--- a/unsloth/kernels/moe/grouped_gemm/reference/moe_ops.py
+++ b/unsloth/kernels/moe/grouped_gemm/reference/moe_ops.py
@@ -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
diff --git a/unsloth/kernels/moe/tests/moe_utils.py b/unsloth/kernels/moe/tests/moe_utils.py
index 17ad797104..18a30dd349 100644
--- a/unsloth/kernels/moe/tests/moe_utils.py
+++ b/unsloth/kernels/moe/tests/moe_utils.py
@@ -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(
diff --git a/unsloth/kernels/moe/tests/test_llama4_moe.py b/unsloth/kernels/moe/tests/test_llama4_moe.py
new file mode 100644
index 0000000000..16d1611dd7
--- /dev/null
+++ b/unsloth/kernels/moe/tests/test_llama4_moe.py
@@ -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,
+ )
diff --git a/unsloth/kernels/moe/tests/test_qwen3_moe.py b/unsloth/kernels/moe/tests/test_qwen3_moe.py
index 8973eb0a4d..ee4277ce23 100644
--- a/unsloth/kernels/moe/tests/test_qwen3_moe.py
+++ b/unsloth/kernels/moe/tests/test_qwen3_moe.py
@@ -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)