From 15bce315ac537e22886e1ebfbfd0dfa0f6b7ee9b Mon Sep 17 00:00:00 2001 From: jeromeku Date: Wed, 28 May 2025 03:26:35 -0700 Subject: [PATCH 1/3] Llama4 MoE Grouped GEMM (#2639) * add llama4 reference layer * add llama4 reference impl * formatting --- unsloth/kernels/moe/README.md | 23 +- .../moe/benchmark/benchmark_fused_moe.py | 422 ++++++++++------- unsloth/kernels/moe/benchmark/utils.py | 93 +++- .../reference/layers/llama4_moe.py | 434 ++++++++++++++++++ .../reference/layers/qwen3_moe.py | 345 ++++++++++++++ .../moe/grouped_gemm/reference/moe_ops.py | 202 +------- unsloth/kernels/moe/tests/moe_utils.py | 5 +- unsloth/kernels/moe/tests/test_llama4_moe.py | 259 +++++++++++ unsloth/kernels/moe/tests/test_qwen3_moe.py | 29 +- 9 files changed, 1429 insertions(+), 383 deletions(-) create mode 100644 unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py create mode 100644 unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py create mode 100644 unsloth/kernels/moe/tests/test_llama4_moe.py 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) From e2f3a6567d8f07068ab30b14b1542bfc41b1f0a0 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 28 May 2025 03:27:48 -0700 Subject: [PATCH 2/3] Create LICENSE --- unsloth/kernels/moe/grouped_gemm/LICENSE | 661 +++++++++++++++++++++++ 1 file changed, 661 insertions(+) create mode 100644 unsloth/kernels/moe/grouped_gemm/LICENSE 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 From 5d90c8303f90ad0d72a7648b8e9eabd04b9da6ad Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 28 May 2025 06:15:12 -0700 Subject: [PATCH 3/3] Latest TRL, GRPO + Bug fixes (#2645) * Update vision.py * Update vision.py * Update vision.py * Update vision.py * model_type_arch * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * check * Update _utils.py * Update loader.py * Update loader.py * Remove prints * Update README.md typo * Update _utils.py * Update _utils.py * versioning * Update _utils.py * Update _utils.py * Update _utils.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update vision.py * HF Transfer * fix(utils): add missing importlib import to fix NameError (#2134) This commit fixes a NameError that occurs when `importlib` is referenced in _utils.py without being imported, especially when UNSLOTH_USE_MODELSCOPE=1 is enabled. By adding the missing import statement, the code will no longer throw a NameError. * Add QLoRA Train and Merge16bit Test (#2130) * add reference and unsloth lora merging tests * add test / dataset printing to test scripts * allow running tests from repo root * add qlora test readme * more readme edits * ruff formatting * additional readme comments * forgot to add actual tests * add apache license * Update pyproject.toml * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Revert * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Bug fix * Update mapper.py * check SDPA for Mistral 3, Pixtral * Update vision.py * Versioning * Update rl_replacements.py * Update README.md * add model registry * move hf hub utils to unsloth/utils * refactor global model info dicts to dataclasses * fix dataclass init * fix llama registration * remove deprecated key function * start registry reog * add llama vision * quant types -> Enum * remap literal quant types to QuantType Enum * add llama model registration * fix quant tag mapping * add qwen2.5 models to registry * add option to include original model in registry * handle quant types per model size * separate registration of base and instruct llama3.2 * add QwenQVQ to registry * add gemma3 to registry * add phi * add deepseek v3 * add deepseek r1 base * add deepseek r1 zero * add deepseek distill llama * add deepseek distill models * remove redundant code when constructing model names * add mistral small to registry * rename model registration methods * rename deepseek registration methods * refactor naming for mistral and phi * add global register models * refactor model registration tests for new registry apis * add model search method * remove deprecated registration api * add quant type test * add registry readme * make llama registration more specific * clear registry when executing individual model registration file * more registry readme updates * Update _auto_install.py * Llama4 * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Synthetic data * Update mapper.py * Xet and Synthetic * Update synthetic.py * Update loader.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update _utils.py * Update pyproject.toml * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update chat_templates.py * Seasame force float16 / float32 * Fix Seasame * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * is_multimodal * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * UNSLOTH_DISABLE_STATIC_GENERATION * Update vision.py * Auto vision detection * Sesame * Whisper * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update _utils.py * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * logging * Update pyproject.toml * Update rl.py --------- Co-authored-by: Jack Shi Wei Lun <87535974+jackswl@users.noreply.github.com> Co-authored-by: naliazheli Co-authored-by: jeromeku Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> --- pyproject.toml | 12 +++++------ unsloth/models/_utils.py | 2 +- unsloth/models/rl.py | 35 ++++++++++++++++++++++++++++++- unsloth/models/rl_replacements.py | 16 +++++++++++++- 4 files changed, 56 insertions(+), 9 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index 523794ee1a..f9a33a861a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -37,10 +37,10 @@ triton = [ ] huggingface = [ - "unsloth_zoo>=2025.5.8", + "unsloth_zoo>=2025.5.10", "packaging", "tyro", - "transformers==4.51.3,!=4.47.0", + "transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2", "datasets>=3.4.1", "sentencepiece>=0.2.0", "tqdm", @@ -48,7 +48,7 @@ huggingface = [ "wheel>=0.42.0", "numpy", "accelerate>=0.34.1", - "trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0,<=0.15.2", + "trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0", "peft>=0.7.1,!=0.11.0", "protobuf<4.0.0", "huggingface_hub", @@ -381,10 +381,10 @@ colab-ampere-torch220 = [ "flash-attn>=2.6.3", ] colab-new = [ - "unsloth_zoo>=2025.5.8", + "unsloth_zoo>=2025.5.9", "packaging", "tyro", - "transformers==4.51.3,!=4.47.0", + "transformers>=4.51.3,!=4.47.0,!=4.52.0,!=4.52.1,!=4.52.2", "datasets>=3.4.1", "sentencepiece>=0.2.0", "tqdm", @@ -399,7 +399,7 @@ colab-new = [ ] colab-no-deps = [ "accelerate>=0.34.1", - "trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0,<=0.15.2", + "trl>=0.7.9,!=0.9.0,!=0.9.1,!=0.9.2,!=0.9.3,!=0.15.0", "peft>=0.7.1", "xformers", "bitsandbytes>=0.45.5", diff --git a/unsloth/models/_utils.py b/unsloth/models/_utils.py index 964e874c58..9325428060 100644 --- a/unsloth/models/_utils.py +++ b/unsloth/models/_utils.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -__version__ = "2025.5.7" +__version__ = "2025.5.8" __all__ = [ "SUPPORTS_BFLOAT16", diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index b385dba2eb..e5cb226433 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -395,7 +395,7 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): if trainer_file in RL_METRICS_CHANGES: process_extra_args = RL_METRICS_CHANGES[trainer_file] for process_extra_arg in process_extra_args: - other_metrics_processor += process_extra_arg(call_args, extra_args) + other_metrics_processor += process_extra_arg(old_RLTrainer_source, old_RLConfig_source) pass # Add statistics as well! @@ -481,6 +481,39 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): extra_args += num_proc_check pass + # Check for loss_type = dr_grpo and scale_rewards for GRPO + if "loss_type" in call_args and "scale_rewards" in call_args: + check_dr_grpo = \ + "if loss_type.lower() == 'dr_grpo':\n"\ + " loss_type = 'dr_grpo'\n"\ + "elif loss_type.lower() == 'dapo':\n"\ + " loss_type = 'dapo'\n"\ + "if loss_type.lower() == 'dr_grpo':\n"\ + " if scale_rewards == None:\n"\ + " scale_rewards = True\n"\ + " elif scale_rewards == True:\n"\ + " print('The Dr GRPO paper recommends setting `scale_rewards` to False! Will override. Set it to `None` to force False.')\n"\ + " scale_rewards = False\n"\ + "elif loss_type.lower() == 'dapo':\n"\ + " print('The DAPO paper recommends `mask_truncated_completions = True`')\n"\ + " print('The DAPO paper recommends `epsilon_high = 0.28`')\n"\ + " mask_truncated_completions = True\n"\ + " epsilon_high = 0.28\n"\ + "\n" + extra_args += check_dr_grpo + pass + + # Check GRPO num_generations mismatch + if "per_device_train_batch_size" in call_args and "num_generations" in call_args: + check_num_generations = \ + "if (per_device_train_batch_size // num_generations) * num_generations != per_device_train_batch_size:\n"\ + " print('Unsloth: We now expect `per_device_train_batch_size` to be a multiple of `num_generations`.\\n"\ + "We will change the batch size of ' + str(per_device_train_batch_size) + ' to the `num_generations` of ' + str(num_generations))\n"\ + " per_device_train_batch_size = num_generations\n"\ + "\n" + extra_args += check_num_generations + pass + # Edit config with anything extra if trainer_file in RL_CONFIG_CHANGES: process_extra_args = RL_CONFIG_CHANGES[trainer_file] diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 2ff0e253e3..171e75d197 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -363,13 +363,27 @@ RL_CONFIG_CHANGES["grpo_trainer"].append(grpo_trainer_fix_batch_size) def grpo_trainer_metrics(RLTrainer_source, RLConfig_source): if "reward_funcs" not in RLTrainer_source: return "" + # For new TRL we have /mean and /std + use_mean = "rewards/{reward_func_name}/mean" in RLTrainer_source + use_std = "rewards/{reward_func_name}/std" in RLTrainer_source + if not use_mean: + use_normal = "rewards/{reward_func_name}" in RLTrainer_source + else: + use_normal = False + pass + log_metrics = \ "if not isinstance(reward_funcs, list): _reward_funcs = [reward_funcs]\n"\ "else: _reward_funcs = reward_funcs\n"\ "for reward_func in _reward_funcs:\n"\ " try:\n"\ " reward_func_name = reward_func.__name__\n"\ - " other_metrics.append(f'rewards/{reward_func_name}')\n"\ + f" if {use_mean}:\n"\ + " other_metrics.append(f'rewards/{reward_func_name}/mean')\n"\ + f" if {use_std}:\n"\ + " other_metrics.append(f'rewards/{reward_func_name}/std')\n"\ + f" if {use_normal}:\n"\ + " other_metrics.append(f'rewards/{reward_func_name}')\n"\ " except: pass\n" return log_metrics pass