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/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)