Merge branch 'main' into nightly
This commit is contained in:
commit
5f444abcd2
10 changed files with 2079 additions and 372 deletions
|
|
@ -43,18 +43,26 @@ sets of tests. E.g., to run forward tests with autotune turned on: `pytest -sv
|
|||
- `grouped_gemm/tests/test_qwen3_moe.py`: end to end test for Qwen3 MoE block. IMPORTANT: read `tests/run_qwen3_moe_tests.sh` as well as notes in the test itself for complications when running parametrized pytest test suites and triton / autotune. TLDR: use the test script and NOT pytest to run the tests.
|
||||
|
||||
### Benchmarks
|
||||
- `grouped_gemm/benchmark/benchmark_fused_moe.py`: benchmarks HF `Qwen3SpareMOEBlock` against the fused implementation
|
||||
- `grouped_gemm/benchmark/benchmark_fused_moe.py`: benchmarks HF `Qwen3SpareMOEBlock` or `Llama4TextMoe` against the fused implementation
|
||||
|
||||
|
||||
Running with these flags on an `H100` to bench forward pass (run with `--help` to see all available flags):
|
||||
|
||||
For `Qwen3-30B-A3B`:
|
||||
```
|
||||
python benchmark/benchmark_fused_moe.py --mode forward --seqlen 1024 --permute_x --permute_y --autotune
|
||||
python benchmark/benchmark_fused_moe.py --model qwen3 --mode forward --seqlen 1024 --permute_x --permute_y --autotune
|
||||
```
|
||||
|
||||
For the backward bench:
|
||||
```
|
||||
python benchmark/benchmark_fused_moe.py --mode backward --seqlen 1024 --permute_x --permute_y --autotune
|
||||
python benchmark/benchmark_fused_moe.py --model qwen3 --mode backward --seqlen 1024 --permute_x --permute_y --autotune
|
||||
```
|
||||
|
||||
On my machine and env, I get speedups > 25x and 14x respectively.
|
||||
For `Llama-4-Scout-17B-16E`:
|
||||
```
|
||||
python benchmark/benchmark_fused_moe.py --model llama4 --autotune --mode=forward --permute_y
|
||||
```
|
||||
Ditto for backwards.
|
||||
|
||||
### Notes
|
||||
- Tested and benched on `H100`, though should run on Ampere and possibly even earlier gpu generations though the autotuning configs will need to be adjusted.
|
||||
|
|
@ -62,6 +70,7 @@ On my machine and env, I get speedups > 25x and 14x respectively.
|
|||
- The kernels can be run either as autotuned (see `autotuning.py`) or with manually specified config (see `tuning.py`). Recommended to run using autotuner since the MoE block requires 2 configs for the forward (2 grouped gemms) and 4 for the backwards (dX and dW per grouped gemm, 2 grouped gemms).
|
||||
- Running with autotuning turned off with the default manual kernel config will result is **highly** sub-optimal performance as it is only meant for testing / debugging purposes.
|
||||
- I've tried to strike a balance between compilation time and autotuning search space -- can probably squeeze even more performance for specific workloads.
|
||||
- The Llama4 reference layer is still highly under-optimized as there are many low-hanging opportunities for further speedups around routing and shared expert calculation.
|
||||
|
||||
TODO:
|
||||
- TMA store: implemented but not enabled currently due to non-determinism arising from triton pipelining bug.
|
||||
|
|
@ -69,4 +78,8 @@ TODO:
|
|||
- Additional optimizations:
|
||||
- Fused / optimized implementations of routing, token sorting, etc.
|
||||
- Better software pipelining within grouped gemm
|
||||
- Threadblock swizzling for better L2 caching
|
||||
- Threadblock swizzling for better L2 caching
|
||||
- Llama4
|
||||
- Fused gather / topk weight merging
|
||||
- Custom topk, gather indices kernel
|
||||
- Shared expert fusion with experts calculation
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
661
unsloth/kernels/moe/grouped_gemm/LICENSE
Normal file
661
unsloth/kernels/moe/grouped_gemm/LICENSE
Normal file
|
|
@ -0,0 +1,661 @@
|
|||
GNU AFFERO GENERAL PUBLIC LICENSE
|
||||
Version 3, 19 November 2007
|
||||
|
||||
Copyright (C) 2007 Free Software Foundation, Inc. <https://fsf.org/>
|
||||
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.
|
||||
|
||||
<one line to give the program's name and a brief idea of what it does.>
|
||||
Copyright (C) <year> <name of author>
|
||||
|
||||
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 <https://www.gnu.org/licenses/>.
|
||||
|
||||
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
|
||||
<https://www.gnu.org/licenses/>.
|
||||
434
unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py
Normal file
434
unsloth/kernels/moe/grouped_gemm/reference/layers/llama4_moe.py
Normal file
|
|
@ -0,0 +1,434 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers.models.llama4 import Llama4TextConfig
|
||||
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
|
||||
|
||||
from grouped_gemm.interface import grouped_gemm
|
||||
from grouped_gemm.kernels.tuning import (
|
||||
KernelConfigBackward_dW,
|
||||
KernelConfigBackward_dX,
|
||||
KernelConfigForward,
|
||||
)
|
||||
from grouped_gemm.reference.moe_ops import (
|
||||
get_routing_indices,
|
||||
permute,
|
||||
torch_grouped_gemm,
|
||||
unpermute,
|
||||
)
|
||||
|
||||
"""
|
||||
Reference implementation of Llama4 MoE block using triton grouped gemm.
|
||||
|
||||
`Llama4GroupedGemmTextMoe` is the HF `Llama4TextMoe` block implemented with a torch-native grouped gemm.
|
||||
`Llama4TritonTextMoe` is the HF `Llama4TextMoe` implemented with triton grouped gemm.
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class Llama4MoeResult:
|
||||
token_counts_by_expert: torch.Tensor
|
||||
gather_indices: torch.Tensor
|
||||
topk_weights: torch.Tensor
|
||||
hidden_states_after_weight_merge: torch.Tensor
|
||||
first_gemm: torch.Tensor
|
||||
intermediate: torch.Tensor
|
||||
second_gemm: torch.Tensor
|
||||
hidden_states_unpermute: torch.Tensor
|
||||
shared_expert_out: torch.Tensor
|
||||
final_out: torch.Tensor
|
||||
router_logits: torch.Tensor = None
|
||||
|
||||
|
||||
class Llama4GroupedGemmTextMoe(Llama4TextMoe):
|
||||
EXPERT_WEIGHT_NAMES = ["experts.gate_up_proj", "experts.down_proj"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Llama4TextConfig,
|
||||
overlap_router_shared=False,
|
||||
verbose=False,
|
||||
debug=False,
|
||||
):
|
||||
super().__init__(config)
|
||||
self.overlap_router_shared = overlap_router_shared
|
||||
self.verbose = verbose
|
||||
self.debug = debug
|
||||
|
||||
# Permute in-place expert weights
|
||||
E, K, N = self.num_experts, self.hidden_dim, self.experts.expert_dim
|
||||
assert self.experts.gate_up_proj.shape == torch.Size([E, K, 2 * N]), (
|
||||
f"{self.experts.gate_up_proj.shape} != {[E, K, 2 * N]}"
|
||||
)
|
||||
permuted_shape = [E, 2 * N, K]
|
||||
permuted_stride = [2 * N * K, K, 1]
|
||||
if verbose:
|
||||
print(
|
||||
f"Changing gate_up_proj from {self.experts.gate_up_proj.size()}:{self.experts.gate_up_proj.stride()} to {permuted_shape}:{permuted_stride}"
|
||||
)
|
||||
with torch.no_grad():
|
||||
self.experts.gate_up_proj.as_strided_(permuted_shape, permuted_stride)
|
||||
|
||||
if verbose:
|
||||
print(
|
||||
f"{self.experts.gate_up_proj.shape}:{self.experts.gate_up_proj.stride()}"
|
||||
)
|
||||
|
||||
assert self.experts.down_proj.shape == torch.Size([E, N, K]), (
|
||||
f"{self.experts.down_proj.shape} != {[E, N, K]}"
|
||||
)
|
||||
permuted_shape = [E, K, N]
|
||||
permuted_stride = [K * N, N, 1]
|
||||
if verbose:
|
||||
print(
|
||||
f"Changing down_proj from {self.experts.down_proj.size()}:{self.experts.down_proj.stride()} to {permuted_shape}:{permuted_stride}"
|
||||
)
|
||||
|
||||
with torch.no_grad():
|
||||
self.experts.down_proj.as_strided_(permuted_shape, permuted_stride)
|
||||
|
||||
if verbose:
|
||||
print(f"{self.experts.down_proj.shape}:{self.experts.down_proj.stride()}")
|
||||
|
||||
if overlap_router_shared:
|
||||
self.shared_expert_stream = torch.cuda.Stream()
|
||||
self.default_event = torch.cuda.Event()
|
||||
self.shared_expert_end_event = torch.cuda.Event()
|
||||
|
||||
@torch.no_grad
|
||||
def copy_weights(self, other: Llama4TextMoe):
|
||||
for name, param_to_copy in other.named_parameters():
|
||||
if self.verbose:
|
||||
print(f"Copying {name} with shape {param_to_copy.shape}")
|
||||
param = self.get_parameter(name)
|
||||
|
||||
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
|
||||
param_to_copy = param_to_copy.permute(0, 2, 1)
|
||||
|
||||
assert param.shape == param_to_copy.shape, (
|
||||
f"{param.shape} != {param_to_copy.shape}"
|
||||
)
|
||||
param.copy_(param_to_copy)
|
||||
|
||||
return self
|
||||
|
||||
def check_weights(self, other: Llama4TextMoe):
|
||||
for name, other_param in other.named_parameters():
|
||||
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
|
||||
other_param = other_param.permute(0, 2, 1)
|
||||
param = self.get_parameter(name)
|
||||
assert param.equal(other_param), f"Param {name} not equal!"
|
||||
assert param.is_contiguous(), f"{name} not contiguous!"
|
||||
|
||||
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
|
||||
assert x.shape[-1] == 2 * self.experts.expert_dim
|
||||
gate_proj = x[..., : self.experts.expert_dim]
|
||||
up_proj = x[..., self.experts.expert_dim :]
|
||||
return self.experts.act_fn(gate_proj) * up_proj
|
||||
|
||||
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
# router_logits: (batch * sequence_length, n_experts)
|
||||
hidden_states = hidden_states.view(-1, self.hidden_dim)
|
||||
router_logits = self.router(hidden_states)
|
||||
routing_weights, selected_experts = torch.topk(
|
||||
router_logits, self.top_k, dim=-1
|
||||
)
|
||||
|
||||
routing_weights = F.sigmoid(routing_weights.float()).to(hidden_states.dtype)
|
||||
|
||||
return router_logits, routing_weights, selected_experts
|
||||
|
||||
def get_token_counts_and_gather_indices(
|
||||
self, selected_experts: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
token_counts_by_expert, gather_indices = get_routing_indices(
|
||||
selected_experts, self.num_experts
|
||||
)
|
||||
assert not token_counts_by_expert.requires_grad
|
||||
assert not gather_indices.requires_grad
|
||||
return token_counts_by_expert, gather_indices
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
""" """
|
||||
batch_size, sequence_length, hidden_dim = hidden_states.shape
|
||||
num_tokens = batch_size * sequence_length
|
||||
total_tokens = num_tokens * self.top_k
|
||||
hidden_states = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
if self.overlap_router_shared:
|
||||
# Marker for all prior ops on default stream
|
||||
self.default_event.record()
|
||||
|
||||
router_logits, routing_weights, selected_experts = self.run_router(
|
||||
hidden_states
|
||||
)
|
||||
assert routing_weights.shape == (num_tokens, self.top_k), (
|
||||
f"{routing_weights.shape} != {(num_tokens, self.top_k)}"
|
||||
)
|
||||
|
||||
if self.overlap_router_shared:
|
||||
with torch.cuda.stream(self.shared_expert_stream):
|
||||
# Ensure prior kernels on default stream complete
|
||||
self.default_event.wait()
|
||||
|
||||
shared_expert_out = self.shared_expert(hidden_states)
|
||||
# Ensure hidden states remains valid on this stream
|
||||
hidden_states.record_stream(self.shared_expert_stream)
|
||||
|
||||
self.shared_expert_end_event.record()
|
||||
|
||||
# Ensure shared expert still valid on default stream
|
||||
shared_expert_out.record_stream(torch.cuda.current_stream())
|
||||
self.shared_expert_end_event.wait()
|
||||
else:
|
||||
shared_expert_out = self.shared_expert(hidden_states)
|
||||
|
||||
hidden_states = (
|
||||
hidden_states.view(num_tokens, self.top_k, hidden_dim)
|
||||
* routing_weights[..., None]
|
||||
)
|
||||
|
||||
if self.top_k > 1:
|
||||
hidden_states = hidden_states.sum(dim=1)
|
||||
hidden_states_after_weight_merge = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
|
||||
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
|
||||
token_counts_by_expert, gather_indices = (
|
||||
self.get_token_counts_and_gather_indices(selected_experts)
|
||||
)
|
||||
|
||||
# 2. Permute tokens from token order to expert order
|
||||
hidden_states = permute(
|
||||
hidden_states_after_weight_merge, gather_indices, self.top_k
|
||||
)
|
||||
assert hidden_states.shape == (total_tokens, hidden_dim)
|
||||
|
||||
# Start expert computation
|
||||
first_gemm = torch_grouped_gemm(
|
||||
X=hidden_states, W=self.experts.gate_up_proj, m_sizes=token_counts_by_expert
|
||||
)
|
||||
assert first_gemm.shape == (total_tokens, 2 * self.experts.expert_dim)
|
||||
|
||||
intermediate = self.act_and_mul(first_gemm)
|
||||
assert intermediate.shape == (total_tokens, self.experts.expert_dim)
|
||||
|
||||
# See comment above
|
||||
second_gemm = torch_grouped_gemm(
|
||||
X=intermediate, W=self.experts.down_proj, m_sizes=token_counts_by_expert
|
||||
)
|
||||
assert second_gemm.shape == (total_tokens, hidden_dim)
|
||||
|
||||
# Post-processing
|
||||
hidden_states_unpermute = unpermute(second_gemm, gather_indices)
|
||||
assert hidden_states_unpermute.shape == (total_tokens, hidden_dim)
|
||||
# grouped_gemm_out = hidden_states.view(batch_size, sequence_length, hidden_dim)
|
||||
|
||||
final_out = hidden_states_unpermute + shared_expert_out
|
||||
|
||||
result = (
|
||||
Llama4MoeResult(
|
||||
token_counts_by_expert=token_counts_by_expert,
|
||||
gather_indices=gather_indices,
|
||||
topk_weights=routing_weights,
|
||||
hidden_states_after_weight_merge=hidden_states_after_weight_merge,
|
||||
first_gemm=first_gemm,
|
||||
intermediate=intermediate,
|
||||
second_gemm=second_gemm,
|
||||
hidden_states_unpermute=hidden_states_unpermute,
|
||||
shared_expert_out=shared_expert_out,
|
||||
final_out=final_out,
|
||||
router_logits=router_logits,
|
||||
)
|
||||
if self.debug
|
||||
else (final_out, routing_weights)
|
||||
)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class Llama4TritonTextMoe(Llama4GroupedGemmTextMoe):
|
||||
def __init__(
|
||||
self,
|
||||
config: Llama4TextConfig,
|
||||
overlap_router_shared=False,
|
||||
permute_x: bool = False,
|
||||
permute_y: bool = True,
|
||||
autotune: bool = True,
|
||||
kernel_config_fwd: KernelConfigForward = None,
|
||||
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
|
||||
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
|
||||
dW_only: bool = False,
|
||||
dX_only: bool = False,
|
||||
verbose=False,
|
||||
):
|
||||
super().__init__(config, overlap_router_shared=overlap_router_shared)
|
||||
assert not permute_x, (
|
||||
"Llama4 triton grouped gemm does not support permute x due to pre-multiplication of router weights"
|
||||
)
|
||||
self.permute_x = permute_x
|
||||
self.permute_y = permute_y
|
||||
self.autotune = autotune
|
||||
if not autotune:
|
||||
assert (
|
||||
kernel_config_fwd is not None
|
||||
and kernel_config_bwd_dW is not None
|
||||
and kernel_config_bwd_dX is not None
|
||||
), "Kernel configs must be provided if autotune is False"
|
||||
self.kernel_config_fwd = kernel_config_fwd
|
||||
self.kernel_config_bwd_dW = kernel_config_bwd_dW
|
||||
self.kernel_config_bwd_dX = kernel_config_bwd_dX
|
||||
self.dW_only = dW_only
|
||||
self.dX_only = dX_only
|
||||
|
||||
@torch.no_grad
|
||||
def copy_weights(self, other: Llama4TextMoe):
|
||||
for name, param_to_copy in other.named_parameters():
|
||||
if self.verbose:
|
||||
print(f"Copying {name} with shape {param_to_copy.shape}")
|
||||
param = self.get_parameter(name)
|
||||
|
||||
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
|
||||
param_to_copy = param_to_copy.permute(0, 2, 1)
|
||||
|
||||
assert param.shape == param_to_copy.shape, (
|
||||
f"{param.shape} != {param_to_copy.shape}"
|
||||
)
|
||||
param.copy_(param_to_copy)
|
||||
|
||||
return self
|
||||
|
||||
def check_weights(self, other: Llama4TextMoe):
|
||||
for name, other_param in other.named_parameters():
|
||||
if any(n in name for n in self.EXPERT_WEIGHT_NAMES):
|
||||
other_param = other_param.permute(0, 2, 1)
|
||||
param = self.get_parameter(name)
|
||||
assert param.equal(other_param), f"Param {name} not equal!"
|
||||
assert param.is_contiguous(), f"{name} not contiguous!"
|
||||
|
||||
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
|
||||
assert x.shape[-1] == 2 * self.experts.expert_dim
|
||||
gate_proj = x[..., : self.experts.expert_dim]
|
||||
up_proj = x[..., self.experts.expert_dim :]
|
||||
return self.experts.act_fn(gate_proj) * up_proj
|
||||
|
||||
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
# router_logits: (batch * sequence_length, n_experts)
|
||||
hidden_states = hidden_states.view(-1, self.hidden_dim)
|
||||
router_logits = self.router(hidden_states)
|
||||
routing_weights, selected_experts = torch.topk(
|
||||
router_logits, self.top_k, dim=-1
|
||||
)
|
||||
|
||||
routing_weights = F.sigmoid(routing_weights.float()).to(hidden_states.dtype)
|
||||
|
||||
return router_logits, routing_weights, selected_experts
|
||||
|
||||
def get_token_counts_and_gather_indices(
|
||||
self, selected_experts: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
token_counts_by_expert, gather_indices = get_routing_indices(
|
||||
selected_experts, self.num_experts
|
||||
)
|
||||
assert not token_counts_by_expert.requires_grad
|
||||
assert not gather_indices.requires_grad
|
||||
return token_counts_by_expert, gather_indices
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
""" """
|
||||
batch_size, sequence_length, hidden_dim = hidden_states.shape
|
||||
num_tokens = batch_size * sequence_length
|
||||
total_tokens = num_tokens * self.top_k
|
||||
hidden_states = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
if self.overlap_router_shared:
|
||||
# Marker for all prior ops on default stream
|
||||
self.default_event.record()
|
||||
|
||||
router_logits, routing_weights, selected_experts = self.run_router(
|
||||
hidden_states
|
||||
)
|
||||
assert routing_weights.shape == (num_tokens, self.top_k), (
|
||||
f"{routing_weights.shape} != {(num_tokens, self.top_k)}"
|
||||
)
|
||||
|
||||
if self.overlap_router_shared:
|
||||
with torch.cuda.stream(self.shared_expert_stream):
|
||||
# Ensure prior kernels on default stream complete
|
||||
self.default_event.wait()
|
||||
|
||||
shared_expert_out = self.shared_expert(hidden_states)
|
||||
# Ensure hidden states remains valid on this stream
|
||||
hidden_states.record_stream(self.shared_expert_stream)
|
||||
|
||||
self.shared_expert_end_event.record()
|
||||
|
||||
# Ensure shared expert still valid on default stream
|
||||
shared_expert_out.record_stream(torch.cuda.current_stream())
|
||||
self.shared_expert_end_event.wait()
|
||||
else:
|
||||
shared_expert_out = self.shared_expert(hidden_states)
|
||||
|
||||
hidden_states = (
|
||||
hidden_states.view(num_tokens, self.top_k, hidden_dim)
|
||||
* routing_weights[..., None]
|
||||
)
|
||||
|
||||
if self.top_k > 1:
|
||||
hidden_states = hidden_states.sum(dim=1)
|
||||
hidden_states = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
|
||||
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
|
||||
token_counts_by_expert, gather_indices = (
|
||||
self.get_token_counts_and_gather_indices(selected_experts)
|
||||
)
|
||||
|
||||
# 2. Permute tokens from token order to expert order
|
||||
hidden_states = permute(hidden_states, gather_indices, self.top_k)
|
||||
assert hidden_states.shape == (total_tokens, hidden_dim)
|
||||
|
||||
# Start expert computation
|
||||
hidden_states = grouped_gemm(
|
||||
X=hidden_states,
|
||||
W=self.experts.gate_up_proj,
|
||||
m_sizes=token_counts_by_expert,
|
||||
gather_indices=gather_indices,
|
||||
topk=self.top_k,
|
||||
permute_x=self.permute_x,
|
||||
permute_y=False, # output of first grouped gemm should never be permuted
|
||||
autotune=self.autotune,
|
||||
kernel_config_fwd=self.kernel_config_fwd,
|
||||
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
|
||||
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
|
||||
is_first_gemm=True,
|
||||
dW_only=self.dW_only,
|
||||
dX_only=self.dX_only,
|
||||
)
|
||||
hidden_states = self.act_and_mul(hidden_states)
|
||||
hidden_states = grouped_gemm(
|
||||
X=hidden_states,
|
||||
W=self.experts.down_proj,
|
||||
m_sizes=token_counts_by_expert,
|
||||
gather_indices=gather_indices,
|
||||
topk=self.top_k,
|
||||
permute_x=False,
|
||||
permute_y=self.permute_y,
|
||||
autotune=self.autotune,
|
||||
kernel_config_fwd=self.kernel_config_fwd,
|
||||
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
|
||||
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
|
||||
is_first_gemm=False,
|
||||
dW_only=self.dW_only,
|
||||
dX_only=self.dX_only,
|
||||
)
|
||||
|
||||
# Post-processing
|
||||
# 1. Unpermute from expert order to token order
|
||||
if not self.permute_y:
|
||||
hidden_states = unpermute(hidden_states, gather_indices)
|
||||
hidden_states += shared_expert_out
|
||||
|
||||
return hidden_states, routing_weights
|
||||
345
unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py
Normal file
345
unsloth/kernels/moe/grouped_gemm/reference/layers/qwen3_moe.py
Normal file
|
|
@ -0,0 +1,345 @@
|
|||
from dataclasses import dataclass
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from transformers.models.qwen3_moe.configuration_qwen3_moe import Qwen3MoeConfig
|
||||
from transformers.models.qwen3_moe.modeling_qwen3_moe import (
|
||||
ACT2FN,
|
||||
Qwen3MoeSparseMoeBlock,
|
||||
)
|
||||
|
||||
from grouped_gemm.interface import grouped_gemm
|
||||
from grouped_gemm.kernels.tuning import (
|
||||
KernelConfigBackward_dW,
|
||||
KernelConfigBackward_dX,
|
||||
KernelConfigForward,
|
||||
)
|
||||
from grouped_gemm.reference.moe_ops import (
|
||||
get_routing_indices,
|
||||
permute,
|
||||
torch_grouped_gemm,
|
||||
unpermute,
|
||||
)
|
||||
|
||||
"""
|
||||
Reference implementation of HF Qwen3 MoE block using grouped gemm.
|
||||
|
||||
The Qwen3MoeGroupedGEMMBlock is a reference torch-native implemention.
|
||||
Qwen3MoeFusedGroupedGEMMBlock is a version using the triton grouped gemm kernel.
|
||||
|
||||
NOTE: This is NOT to be used for production as it contains many extra checks and saves all intermediate results for debugging.
|
||||
"""
|
||||
|
||||
|
||||
@dataclass
|
||||
class GroupedGEMMResult:
|
||||
token_counts_by_expert: torch.Tensor
|
||||
gather_indices: torch.Tensor
|
||||
topk_weights: torch.Tensor
|
||||
first_gemm: torch.Tensor
|
||||
intermediate: torch.Tensor
|
||||
second_gemm: torch.Tensor
|
||||
hidden_states_unpermute: torch.Tensor
|
||||
hidden_states: torch.Tensor # final output
|
||||
|
||||
|
||||
class Qwen3MoeGroupedGEMMBlock(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
gate: torch.Tensor,
|
||||
gate_up_proj: torch.Tensor,
|
||||
down_proj: torch.Tensor,
|
||||
):
|
||||
super().__init__()
|
||||
self.num_experts = config.num_experts
|
||||
self.top_k = config.num_experts_per_tok
|
||||
self.norm_topk_prob = config.norm_topk_prob
|
||||
self.hidden_size = config.hidden_size
|
||||
self.moe_intermediate_size = config.moe_intermediate_size
|
||||
|
||||
assert gate.shape == (config.num_experts, config.hidden_size)
|
||||
assert gate_up_proj.shape == (
|
||||
config.num_experts,
|
||||
2 * config.moe_intermediate_size,
|
||||
config.hidden_size,
|
||||
)
|
||||
assert down_proj.shape == (
|
||||
config.num_experts,
|
||||
config.hidden_size,
|
||||
config.moe_intermediate_size,
|
||||
)
|
||||
|
||||
# gating
|
||||
self.gate = torch.nn.Parameter(gate)
|
||||
|
||||
# experts
|
||||
self.gate_up_proj = torch.nn.Parameter(gate_up_proj, requires_grad=True)
|
||||
self.down_proj = torch.nn.Parameter(down_proj, requires_grad=True)
|
||||
self.act_fn = ACT2FN[config.hidden_act]
|
||||
|
||||
@staticmethod
|
||||
def extract_hf_weights(moe_block: Qwen3MoeSparseMoeBlock):
|
||||
config: Qwen3MoeConfig = moe_block.experts[0].config
|
||||
num_experts = config.num_experts
|
||||
|
||||
gate = moe_block.gate.weight.data
|
||||
gate_proj = torch.stack(
|
||||
[moe_block.experts[i].gate_proj.weight.data for i in range(num_experts)],
|
||||
dim=0,
|
||||
)
|
||||
up_proj = torch.stack(
|
||||
[moe_block.experts[i].up_proj.weight.data for i in range(num_experts)],
|
||||
dim=0,
|
||||
)
|
||||
down_proj = torch.stack(
|
||||
[moe_block.experts[i].down_proj.weight.data for i in range(num_experts)],
|
||||
dim=0,
|
||||
)
|
||||
gate_up_proj = torch.cat([gate_proj, up_proj], dim=1)
|
||||
return gate, gate_up_proj, down_proj
|
||||
|
||||
@classmethod
|
||||
def from_hf(cls, moe_block: Qwen3MoeSparseMoeBlock):
|
||||
config: Qwen3MoeConfig = moe_block.experts[0].config
|
||||
gate, gate_up_proj, down_proj = cls.extract_hf_weights(moe_block)
|
||||
return cls(config, gate, gate_up_proj, down_proj)
|
||||
|
||||
def check_weights(self, moe_block: Qwen3MoeSparseMoeBlock):
|
||||
for i in range(self.num_experts):
|
||||
assert self.gate_up_proj[i].equal(
|
||||
torch.cat(
|
||||
[
|
||||
moe_block.experts[i].gate_proj.weight.data,
|
||||
moe_block.experts[i].up_proj.weight.data,
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
)
|
||||
assert self.down_proj[i].equal(moe_block.experts[i].down_proj.weight.data)
|
||||
|
||||
def act_and_mul(self, x: torch.Tensor) -> torch.Tensor:
|
||||
assert x.shape[-1] == 2 * self.moe_intermediate_size
|
||||
gate_proj = x[..., : self.moe_intermediate_size]
|
||||
up_proj = x[..., self.moe_intermediate_size :]
|
||||
return self.act_fn(gate_proj) * up_proj
|
||||
|
||||
def run_router(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
# router_logits: (batch * sequence_length, n_experts)
|
||||
router_logits = torch.nn.functional.linear(hidden_states, self.gate)
|
||||
|
||||
routing_weights = F.softmax(router_logits, dim=1, dtype=torch.float)
|
||||
routing_weights, selected_experts = torch.topk(
|
||||
routing_weights, self.top_k, dim=-1
|
||||
)
|
||||
if self.norm_topk_prob: # only diff with mixtral sparse moe block!
|
||||
routing_weights /= routing_weights.sum(dim=-1, keepdim=True)
|
||||
# we cast back to the input dtype
|
||||
routing_weights = routing_weights.to(hidden_states.dtype)
|
||||
|
||||
return router_logits, routing_weights, selected_experts
|
||||
|
||||
def get_token_counts_and_gather_indices(
|
||||
self, selected_experts: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
token_counts_by_expert, gather_indices = get_routing_indices(
|
||||
selected_experts, self.num_experts
|
||||
)
|
||||
assert not token_counts_by_expert.requires_grad
|
||||
assert not gather_indices.requires_grad
|
||||
return token_counts_by_expert, gather_indices
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
""" """
|
||||
batch_size, sequence_length, hidden_dim = hidden_states.shape
|
||||
num_tokens = batch_size * sequence_length
|
||||
total_tokens = num_tokens * self.top_k
|
||||
|
||||
hidden_states = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
router_logits, routing_weights, selected_experts = self.run_router(
|
||||
hidden_states
|
||||
)
|
||||
|
||||
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
|
||||
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
|
||||
token_counts_by_expert, gather_indices = (
|
||||
self.get_token_counts_and_gather_indices(selected_experts)
|
||||
)
|
||||
|
||||
# 2. Permute tokens from token order to expert order
|
||||
hidden_states = permute(hidden_states, gather_indices, self.top_k)
|
||||
assert hidden_states.shape == (total_tokens, hidden_dim)
|
||||
|
||||
# Start expert computation
|
||||
first_gemm = torch_grouped_gemm(
|
||||
X=hidden_states, W=self.gate_up_proj, m_sizes=token_counts_by_expert
|
||||
)
|
||||
assert first_gemm.shape == (total_tokens, 2 * self.moe_intermediate_size)
|
||||
intermediate = self.act_and_mul(first_gemm)
|
||||
assert intermediate.shape == (total_tokens, self.moe_intermediate_size)
|
||||
second_gemm = torch_grouped_gemm(
|
||||
X=intermediate, W=self.down_proj, m_sizes=token_counts_by_expert
|
||||
)
|
||||
assert second_gemm.shape == (total_tokens, hidden_dim)
|
||||
|
||||
# Post-processing
|
||||
# 1. Unpermute from expert order to token order
|
||||
hidden_states_unpermute = unpermute(second_gemm, gather_indices)
|
||||
assert hidden_states_unpermute.shape == (total_tokens, hidden_dim)
|
||||
|
||||
# 2. Merge topk weights
|
||||
hidden_states = (
|
||||
hidden_states_unpermute.view(num_tokens, self.top_k, hidden_dim)
|
||||
* routing_weights[..., None]
|
||||
)
|
||||
hidden_states = hidden_states.sum(dim=1)
|
||||
assert hidden_states.shape == (num_tokens, hidden_dim)
|
||||
|
||||
hidden_states = hidden_states.view(batch_size, sequence_length, hidden_dim)
|
||||
return GroupedGEMMResult(
|
||||
token_counts_by_expert=token_counts_by_expert,
|
||||
gather_indices=gather_indices,
|
||||
topk_weights=routing_weights,
|
||||
first_gemm=first_gemm,
|
||||
intermediate=intermediate,
|
||||
second_gemm=second_gemm,
|
||||
hidden_states_unpermute=hidden_states_unpermute,
|
||||
hidden_states=hidden_states,
|
||||
), router_logits
|
||||
|
||||
|
||||
class Qwen3MoeFusedGroupedGEMMBlock(Qwen3MoeGroupedGEMMBlock):
|
||||
def __init__(
|
||||
self,
|
||||
config: Qwen3MoeConfig,
|
||||
gate: torch.Tensor,
|
||||
gate_up_proj: torch.Tensor,
|
||||
down_proj: torch.Tensor,
|
||||
permute_x: bool = True,
|
||||
permute_y: bool = True,
|
||||
autotune: bool = True,
|
||||
kernel_config_fwd: KernelConfigForward = None,
|
||||
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
|
||||
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
|
||||
dW_only: bool = False,
|
||||
dX_only: bool = False,
|
||||
):
|
||||
super().__init__(config, gate, gate_up_proj, down_proj)
|
||||
self.permute_x = permute_x
|
||||
self.permute_y = permute_y
|
||||
self.autotune = autotune
|
||||
if not autotune:
|
||||
assert (
|
||||
kernel_config_fwd is not None
|
||||
and kernel_config_bwd_dW is not None
|
||||
and kernel_config_bwd_dX is not None
|
||||
), "Kernel configs must be provided if autotune is False"
|
||||
self.kernel_config_fwd = kernel_config_fwd
|
||||
self.kernel_config_bwd_dW = kernel_config_bwd_dW
|
||||
self.kernel_config_bwd_dX = kernel_config_bwd_dX
|
||||
self.dW_only = dW_only
|
||||
self.dX_only = dX_only
|
||||
|
||||
@classmethod
|
||||
def from_hf(
|
||||
cls,
|
||||
moe_block: Qwen3MoeSparseMoeBlock,
|
||||
permute_x: bool = True,
|
||||
permute_y: bool = True,
|
||||
autotune: bool = True,
|
||||
kernel_config_fwd: KernelConfigForward = None,
|
||||
kernel_config_bwd_dW: KernelConfigBackward_dW = None,
|
||||
kernel_config_bwd_dX: KernelConfigBackward_dX = None,
|
||||
dW_only: bool = False,
|
||||
dX_only: bool = False,
|
||||
):
|
||||
config: Qwen3MoeConfig = moe_block.experts[0].config
|
||||
gate, gate_up_proj, down_proj = Qwen3MoeGroupedGEMMBlock.extract_hf_weights(
|
||||
moe_block
|
||||
)
|
||||
return cls(
|
||||
config,
|
||||
gate,
|
||||
gate_up_proj,
|
||||
down_proj,
|
||||
permute_x=permute_x,
|
||||
permute_y=permute_y,
|
||||
autotune=autotune,
|
||||
kernel_config_fwd=kernel_config_fwd,
|
||||
kernel_config_bwd_dW=kernel_config_bwd_dW,
|
||||
kernel_config_bwd_dX=kernel_config_bwd_dX,
|
||||
dW_only=dW_only,
|
||||
dX_only=dX_only,
|
||||
)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, sequence_length, hidden_dim = hidden_states.shape
|
||||
num_tokens = batch_size * sequence_length
|
||||
total_tokens = num_tokens * self.top_k
|
||||
|
||||
hidden_states = hidden_states.view(-1, hidden_dim)
|
||||
|
||||
router_logits, routing_weights, selected_experts = self.run_router(
|
||||
hidden_states
|
||||
)
|
||||
# Pre-processing
|
||||
# 1. Compute tokens per expert and indices for gathering tokes from token order to expert order
|
||||
# NOTE: these are auxiliary data structs which don't need to be recorded in autograd graph
|
||||
token_counts_by_expert, gather_indices = (
|
||||
self.get_token_counts_and_gather_indices(selected_experts)
|
||||
)
|
||||
|
||||
# 2. permute_x -> permutation will be fused in prologue of first grouped gemm
|
||||
if not self.permute_x:
|
||||
hidden_states = permute(hidden_states, gather_indices, self.top_k)
|
||||
# Start expert computation
|
||||
hidden_states = grouped_gemm(
|
||||
X=hidden_states,
|
||||
W=self.gate_up_proj,
|
||||
m_sizes=token_counts_by_expert,
|
||||
gather_indices=gather_indices,
|
||||
topk=self.top_k,
|
||||
permute_x=self.permute_x,
|
||||
permute_y=False, # output of first grouped gemm should never be permuted
|
||||
autotune=self.autotune,
|
||||
kernel_config_fwd=self.kernel_config_fwd,
|
||||
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
|
||||
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
|
||||
is_first_gemm=True,
|
||||
dW_only=self.dW_only,
|
||||
dX_only=self.dX_only,
|
||||
)
|
||||
hidden_states = self.act_and_mul(hidden_states)
|
||||
hidden_states = grouped_gemm(
|
||||
X=hidden_states,
|
||||
W=self.down_proj,
|
||||
m_sizes=token_counts_by_expert,
|
||||
gather_indices=gather_indices,
|
||||
topk=self.top_k,
|
||||
permute_x=False,
|
||||
permute_y=self.permute_y,
|
||||
autotune=self.autotune,
|
||||
kernel_config_fwd=self.kernel_config_fwd,
|
||||
kernel_config_bwd_dW=self.kernel_config_bwd_dW,
|
||||
kernel_config_bwd_dX=self.kernel_config_bwd_dX,
|
||||
is_first_gemm=False,
|
||||
dW_only=self.dW_only,
|
||||
dX_only=self.dX_only,
|
||||
)
|
||||
|
||||
# Post-processing
|
||||
# 1. Unpermute from expert order to token order
|
||||
if not self.permute_y:
|
||||
hidden_states = unpermute(hidden_states, gather_indices)
|
||||
|
||||
# 2. Merge topk weights
|
||||
hidden_states = (
|
||||
hidden_states.view(num_tokens, self.top_k, hidden_dim)
|
||||
* routing_weights[..., None]
|
||||
)
|
||||
hidden_states = hidden_states.sum(dim=1)
|
||||
|
||||
hidden_states = hidden_states.view(batch_size, sequence_length, hidden_dim)
|
||||
return hidden_states, router_logits
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
259
unsloth/kernels/moe/tests/test_llama4_moe.py
Normal file
259
unsloth/kernels/moe/tests/test_llama4_moe.py
Normal file
|
|
@ -0,0 +1,259 @@
|
|||
import argparse
|
||||
import sys
|
||||
from contextlib import contextmanager
|
||||
from functools import partial
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig
|
||||
from transformers.models.llama4 import Llama4Config, Llama4TextConfig
|
||||
from transformers.models.llama4.modeling_llama4 import Llama4TextMoe
|
||||
|
||||
from grouped_gemm.kernels.tuning import (
|
||||
KernelConfigBackward_dW,
|
||||
KernelConfigBackward_dX,
|
||||
KernelConfigForward,
|
||||
)
|
||||
from grouped_gemm.reference.layers.llama4_moe import (
|
||||
Llama4GroupedGemmTextMoe,
|
||||
Llama4TritonTextMoe,
|
||||
)
|
||||
|
||||
TOLERANCES = {
|
||||
torch.bfloat16: (1e-2, 1e-2),
|
||||
torch.float16: (1e-3, 1e-3),
|
||||
torch.float: (1e-5, 1e-5),
|
||||
}
|
||||
|
||||
LLAMA4_SCOUT_ID = "meta-llama/Llama-4-Scout-17B-16E"
|
||||
SEED = 42
|
||||
SEQ_LENS = [1024]
|
||||
DTYPES = [torch.bfloat16]
|
||||
# Reduce the number of autotuning configs to prevent excessive runtime
|
||||
NUM_AUTOTUNE_CONFIGS = 50
|
||||
|
||||
|
||||
@contextmanager
|
||||
def annotated_context(prelude, epilogue="Passed!", char="-", num_chars=80):
|
||||
print(char * num_chars)
|
||||
print(prelude)
|
||||
yield
|
||||
print(epilogue)
|
||||
print(char * num_chars)
|
||||
|
||||
|
||||
def get_text_config(model_id):
|
||||
config: Llama4Config = AutoConfig.from_pretrained(model_id)
|
||||
return config.text_config
|
||||
|
||||
|
||||
def prep_triton_kernel_traits(autotune):
|
||||
if not autotune:
|
||||
kernel_config_fwd = KernelConfigForward()
|
||||
kernel_config_bwd_dW = KernelConfigBackward_dW()
|
||||
kernel_config_bwd_dX = KernelConfigBackward_dX()
|
||||
else:
|
||||
from grouped_gemm.kernels.backward import (
|
||||
_autotuned_grouped_gemm_dW_kernel,
|
||||
_autotuned_grouped_gemm_dX_kernel,
|
||||
)
|
||||
from grouped_gemm.kernels.forward import _autotuned_grouped_gemm_forward_kernel
|
||||
|
||||
# Hack to reduce number of autotuning configs
|
||||
_autotuned_grouped_gemm_forward_kernel.configs = (
|
||||
_autotuned_grouped_gemm_forward_kernel.configs[:NUM_AUTOTUNE_CONFIGS]
|
||||
)
|
||||
_autotuned_grouped_gemm_dW_kernel.configs = (
|
||||
_autotuned_grouped_gemm_dW_kernel.configs[:NUM_AUTOTUNE_CONFIGS]
|
||||
)
|
||||
_autotuned_grouped_gemm_dX_kernel.configs = (
|
||||
_autotuned_grouped_gemm_dX_kernel.configs[:NUM_AUTOTUNE_CONFIGS]
|
||||
)
|
||||
|
||||
kernel_config_fwd = None
|
||||
kernel_config_bwd_dW = None
|
||||
kernel_config_bwd_dX = None
|
||||
|
||||
return kernel_config_fwd, kernel_config_bwd_dW, kernel_config_bwd_dX
|
||||
|
||||
|
||||
def sparse_to_dense(t: torch.Tensor):
|
||||
t = t.sum(dim=0).view(-1)
|
||||
return t
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def _check_diff(
|
||||
t1: torch.Tensor,
|
||||
t2: torch.Tensor,
|
||||
atol,
|
||||
rtol,
|
||||
precision=".6f",
|
||||
verbose=False,
|
||||
msg="",
|
||||
):
|
||||
t2 = t2.view_as(t1)
|
||||
diff = t1.sub(t2).abs().max().item()
|
||||
if verbose:
|
||||
if msg == "":
|
||||
msg = "diff"
|
||||
print(f"{msg}: {diff:{precision}}")
|
||||
assert torch.allclose(t1, t2, atol=atol, rtol=rtol)
|
||||
|
||||
|
||||
def run_backwards(y: torch.Tensor, grad_output: torch.Tensor, module: torch.nn.Module):
|
||||
y.backward(grad_output)
|
||||
for name, param in module.named_parameters():
|
||||
assert param.grad is not None, f"{name} missing grad!"
|
||||
|
||||
|
||||
def _check_grads(
|
||||
m1: torch.nn.Module,
|
||||
m2: torch.nn.Module,
|
||||
atol,
|
||||
rtol,
|
||||
precision=".6f",
|
||||
verbose=False,
|
||||
msg="",
|
||||
):
|
||||
for name, param in m1.named_parameters():
|
||||
_check_diff(
|
||||
param.grad,
|
||||
m2.get_parameter(name).grad,
|
||||
atol=atol,
|
||||
rtol=rtol,
|
||||
precision=precision,
|
||||
verbose=verbose,
|
||||
msg=f"{msg}:{name}.grad",
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def model_config():
|
||||
return AutoConfig.from_pretrained(LLAMA4_SCOUT_ID).text_config
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"overlap_router_shared",
|
||||
[False, True],
|
||||
ids=lambda x: "overlap_router_shared" if x else "no_overlap",
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"permute_y", [False, True], ids=lambda x: "permute_y" if x else "no_permute_y"
|
||||
)
|
||||
@pytest.mark.parametrize(
|
||||
"permute_x", [False], ids=lambda x: "permute_x" if x else "no_permute_x"
|
||||
) # Llama4 does not support permute_x
|
||||
@pytest.mark.parametrize(
|
||||
"autotune", [True], ids=lambda x: "autotune" if x else "manual"
|
||||
)
|
||||
@pytest.mark.parametrize("seqlen", SEQ_LENS, ids=lambda x: f"seqlen={x}")
|
||||
@pytest.mark.parametrize("dtype", DTYPES, ids=str)
|
||||
def test_llama4_ref(
|
||||
dtype: torch.dtype,
|
||||
seqlen,
|
||||
autotune: bool,
|
||||
permute_x: bool,
|
||||
permute_y: bool,
|
||||
overlap_router_shared: bool,
|
||||
model_config: Llama4TextConfig, # test fixture
|
||||
bs: int = 1,
|
||||
device="cuda",
|
||||
precision=".6f",
|
||||
verbose=False,
|
||||
):
|
||||
torch.manual_seed(
|
||||
SEED
|
||||
) # Should not be needed when running using pytest -- autouse fixture in conftest.py
|
||||
device = "cuda"
|
||||
hidden_dim = model_config.hidden_size
|
||||
atol, rtol = TOLERANCES[dtype]
|
||||
check_diff = partial(
|
||||
_check_diff, atol=atol, rtol=rtol, precision=precision, verbose=verbose
|
||||
)
|
||||
check_grads = partial(
|
||||
_check_grads, atol=atol, rtol=rtol, precision=precision, verbose=verbose
|
||||
)
|
||||
|
||||
# Reference op -- HF
|
||||
llama4_ref = Llama4TextMoe(model_config).to(dtype=dtype, device=device)
|
||||
|
||||
# Torch grouped gemm impl
|
||||
llama4_gg_ref = Llama4GroupedGemmTextMoe(
|
||||
model_config, overlap_router_shared=overlap_router_shared
|
||||
).to(dtype=dtype, device=device)
|
||||
llama4_gg_ref.copy_weights(llama4_ref)
|
||||
llama4_gg_ref.check_weights(llama4_ref)
|
||||
|
||||
x_ref = torch.randn(
|
||||
bs, seqlen, hidden_dim, dtype=dtype, device=device, requires_grad=True
|
||||
)
|
||||
x_torch_gg = x_ref.detach().clone().requires_grad_()
|
||||
x_triton = x_ref.detach().clone().requires_grad_()
|
||||
|
||||
y_ref, routing_ref = llama4_ref(x_ref)
|
||||
y_torch_gg, routing_torch_gg = llama4_gg_ref(x_torch_gg)
|
||||
assert y_ref.shape == y_torch_gg.shape, f"{y_ref.shape} != {y_torch_gg.shape}"
|
||||
with annotated_context("Testing torch grouped gemm Llama4TextMoe"):
|
||||
check_diff(y_ref, y_torch_gg, msg="y_torch_gg")
|
||||
check_diff(
|
||||
sparse_to_dense(routing_ref), routing_torch_gg, msg="routing_torch_gg"
|
||||
)
|
||||
|
||||
kernel_config_fwd, kernel_config_bwd_dW, kernel_config_bwd_dX = (
|
||||
prep_triton_kernel_traits(autotune)
|
||||
)
|
||||
|
||||
llama4_triton = Llama4TritonTextMoe(
|
||||
model_config,
|
||||
overlap_router_shared=overlap_router_shared,
|
||||
permute_x=permute_x,
|
||||
permute_y=permute_y,
|
||||
autotune=autotune,
|
||||
kernel_config_fwd=kernel_config_fwd,
|
||||
kernel_config_bwd_dW=kernel_config_bwd_dW,
|
||||
kernel_config_bwd_dX=kernel_config_bwd_dX,
|
||||
).to(device=device, dtype=dtype)
|
||||
llama4_triton.copy_weights(llama4_ref)
|
||||
llama4_triton.check_weights(llama4_ref)
|
||||
|
||||
y_triton, routing_triton = llama4_triton(x_triton)
|
||||
with annotated_context("Testing triton grouped gemm Llama4TextMoe forward"):
|
||||
check_diff(y_ref, y_triton, msg="y_triton")
|
||||
check_diff(sparse_to_dense(routing_ref), routing_triton, msg="routing_triton")
|
||||
|
||||
ref_grad = torch.randn_like(y_ref)
|
||||
run_backwards(y_ref, ref_grad, llama4_ref)
|
||||
run_backwards(y_torch_gg, ref_grad, llama4_gg_ref)
|
||||
with annotated_context("Testing torch group gemm Llama4TextMoe backward"):
|
||||
check_grads(llama4_ref, llama4_gg_ref, msg="torch_gg")
|
||||
|
||||
run_backwards(y_triton, ref_grad, llama4_triton)
|
||||
with annotated_context("Testing triton group gemm Llama4TextMoe backward"):
|
||||
check_grads(llama4_ref, llama4_triton, msg="triton")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--seqlen", type=int, default=1024)
|
||||
parser.add_argument(
|
||||
"--dtype", type=str, choices=["bfloat16", "float16"], default="bfloat16"
|
||||
)
|
||||
args = parser.parse_args()
|
||||
args.dtype = getattr(torch, args.dtype)
|
||||
args_dict = vars(args)
|
||||
|
||||
model_id = LLAMA4_SCOUT_ID
|
||||
|
||||
text_config: Llama4TextConfig = get_text_config(model_id)
|
||||
for overlap in [False, True]:
|
||||
test_llama4_ref(
|
||||
seqlen=args.seqlen,
|
||||
model_config=text_config,
|
||||
dtype=args.dtype,
|
||||
autotune=True,
|
||||
permute_x=False,
|
||||
permute_y=True,
|
||||
overlap_router_shared=overlap,
|
||||
verbose=True,
|
||||
)
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue