Enable GPT-OSS FlashInfer MXFP4 on SM120 (#32668)
This commit is contained in:
@@ -13,7 +13,7 @@ import torch
|
||||
|
||||
from sglang.test.ci.ci_register import register_cuda_ci
|
||||
|
||||
register_cuda_ci(est_time=120, stage="base-b", runner_config="1-gpu-large")
|
||||
register_cuda_ci(est_time=120, stage="base-b", runner_config="1-gpu-small")
|
||||
|
||||
|
||||
def _random_weights(num_experts: int, hidden: int, intermediate: int):
|
||||
@@ -247,5 +247,191 @@ def test_dsv4_sm120_matches_direct_flashinfer(monkeypatch):
|
||||
assert torch.equal(actual, expected)
|
||||
|
||||
|
||||
def test_gpt_oss_sm120_padding_layout_and_kernel(monkeypatch):
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("CUDA required")
|
||||
if torch.cuda.get_device_capability() != (12, 0):
|
||||
pytest.skip("SM120 required")
|
||||
pytest.importorskip("flashinfer.fused_moe")
|
||||
|
||||
from flashinfer import block_scale_interleave, mxfp8_quantize
|
||||
from flashinfer.fused_moe import cutlass_fused_moe
|
||||
from flashinfer.fused_moe.core import ActivationType
|
||||
|
||||
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass as runner_module
|
||||
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.moe_runner.runner import MoeRunner
|
||||
from sglang.srt.layers.moe.token_dispatcher.standard import StandardDispatchOutput
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||
from sglang.srt.layers.moe.utils import MoeRunnerBackend
|
||||
from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod
|
||||
|
||||
monkeypatch.setattr(
|
||||
runner_module, "use_symmetric_memory", lambda *args, **kwargs: nullcontext()
|
||||
)
|
||||
monkeypatch.setattr(runner_module, "is_allocation_symmetric", lambda: False)
|
||||
monkeypatch.setattr(runner_module, "get_tp_group", lambda: None)
|
||||
|
||||
num_experts, hidden, intermediate = 4, 160, 160
|
||||
padded_hidden = padded_intermediate = 256
|
||||
w13, w2, w13_scale, w2_scale = _random_weights(num_experts, hidden, intermediate)
|
||||
generator = torch.Generator(device="cuda").manual_seed(2)
|
||||
w13_bias = torch.randn(
|
||||
num_experts,
|
||||
2 * intermediate,
|
||||
dtype=torch.bfloat16,
|
||||
device="cuda",
|
||||
generator=generator,
|
||||
)
|
||||
w2_bias = torch.randn(
|
||||
num_experts,
|
||||
hidden,
|
||||
dtype=torch.bfloat16,
|
||||
device="cuda",
|
||||
generator=generator,
|
||||
)
|
||||
layer = SimpleNamespace(
|
||||
w13_weight=torch.nn.Parameter(
|
||||
w13.view(torch.uint8).clone(), requires_grad=False
|
||||
),
|
||||
w2_weight=torch.nn.Parameter(w2.view(torch.uint8).clone(), requires_grad=False),
|
||||
w13_weight_scale=torch.nn.Parameter(
|
||||
w13_scale.view(torch.uint8).clone(), requires_grad=False
|
||||
),
|
||||
w2_weight_scale=torch.nn.Parameter(
|
||||
w2_scale.view(torch.uint8).clone(), requires_grad=False
|
||||
),
|
||||
w13_weight_bias=torch.nn.Parameter(w13_bias.clone(), requires_grad=False),
|
||||
w2_weight_bias=torch.nn.Parameter(w2_bias.clone(), requires_grad=False),
|
||||
num_local_experts=num_experts,
|
||||
moe_tp_size=1,
|
||||
moe_tp_rank=0,
|
||||
moe_ep_size=1,
|
||||
moe_ep_rank=0,
|
||||
)
|
||||
|
||||
method = Mxfp4MoEMethod.__new__(Mxfp4MoEMethod)
|
||||
method._fi_kernel = "cutlass_sm120"
|
||||
method.num_experts = num_experts
|
||||
method.hidden_size = hidden
|
||||
method.intermediate_size_per_partition = intermediate
|
||||
method._padded_hidden = padded_hidden
|
||||
method._padded_intermediate = padded_intermediate
|
||||
config = MoeRunnerConfig(
|
||||
num_experts=num_experts,
|
||||
num_local_experts=num_experts,
|
||||
hidden_size=hidden,
|
||||
intermediate_size_per_partition=intermediate,
|
||||
top_k=4,
|
||||
activation="silu",
|
||||
is_gated=True,
|
||||
gemm1_alpha=1.702,
|
||||
gemm1_clamp_limit=7.0,
|
||||
)
|
||||
method.moe_runner_config = config
|
||||
method.runner = MoeRunner(MoeRunnerBackend.FLASHINFER_MXFP4, config)
|
||||
method._process_weights_for_sm120_cutlass(layer)
|
||||
|
||||
expected_w13 = torch.zeros(
|
||||
num_experts,
|
||||
2 * padded_intermediate,
|
||||
padded_hidden // 2,
|
||||
dtype=torch.uint8,
|
||||
device="cuda",
|
||||
)
|
||||
expected_w13[:, :intermediate, : hidden // 2] = w13[:, 1::2]
|
||||
expected_w13[
|
||||
:, padded_intermediate : padded_intermediate + intermediate, : hidden // 2
|
||||
] = w13[:, 0::2]
|
||||
expected_w13_scale = torch.zeros(
|
||||
num_experts,
|
||||
2 * padded_intermediate,
|
||||
padded_hidden // 32,
|
||||
dtype=torch.uint8,
|
||||
device="cuda",
|
||||
)
|
||||
expected_w13_scale[:, :intermediate, : hidden // 32] = w13_scale.view(torch.uint8)[
|
||||
:, 1::2
|
||||
]
|
||||
expected_w13_scale[
|
||||
:,
|
||||
padded_intermediate : padded_intermediate + intermediate,
|
||||
: hidden // 32,
|
||||
] = w13_scale.view(torch.uint8)[:, 0::2]
|
||||
expected_w13_scale = block_scale_interleave(expected_w13_scale).reshape_as(
|
||||
expected_w13_scale
|
||||
)
|
||||
assert torch.equal(layer.w13_weight, expected_w13)
|
||||
assert torch.equal(layer.w13_weight_scale, expected_w13_scale)
|
||||
assert torch.equal(layer.w13_weight_bias[:, :intermediate], w13_bias[:, 1::2])
|
||||
assert torch.equal(
|
||||
layer.w13_weight_bias[
|
||||
:, padded_intermediate : padded_intermediate + intermediate
|
||||
],
|
||||
w13_bias[:, 0::2],
|
||||
)
|
||||
assert torch.all(layer.swiglu_alpha == 1.702)
|
||||
assert torch.all(layer.swiglu_beta == 1.0)
|
||||
assert torch.all(layer.swiglu_limit == 7.0)
|
||||
assert layer._mxfp4_backend == "flashinfer_cutlass_sm120"
|
||||
|
||||
x = torch.randn(
|
||||
8,
|
||||
hidden,
|
||||
dtype=torch.bfloat16,
|
||||
device="cuda",
|
||||
generator=generator,
|
||||
)
|
||||
logits = torch.randn(
|
||||
8,
|
||||
num_experts,
|
||||
dtype=torch.float32,
|
||||
device="cuda",
|
||||
generator=generator,
|
||||
)
|
||||
topk_weights, topk_ids = torch.topk(torch.softmax(logits, dim=-1), 4, dim=-1)
|
||||
topk_weights /= topk_weights.sum(dim=-1, keepdim=True)
|
||||
dispatch_output = StandardDispatchOutput(
|
||||
x,
|
||||
None,
|
||||
StandardTopKOutput(topk_weights, topk_ids.to(torch.int32), logits),
|
||||
)
|
||||
actual = method._apply_sm120_cutlass(layer, dispatch_output).hidden_states
|
||||
|
||||
x_padded = torch.nn.functional.pad(x, (0, padded_hidden - hidden))
|
||||
x_quant, x_scale = mxfp8_quantize(
|
||||
x_padded, is_sf_swizzled_layout=True, alignment=32
|
||||
)
|
||||
expected = torch.empty(
|
||||
x.shape[0], padded_hidden, dtype=torch.bfloat16, device="cuda"
|
||||
)
|
||||
cutlass_fused_moe(
|
||||
input=x_quant,
|
||||
token_selected_experts=topk_ids.to(torch.int32),
|
||||
token_final_scales=topk_weights,
|
||||
fc1_expert_weights=layer.w13_weight.view(torch.int64),
|
||||
fc2_expert_weights=layer.w2_weight.view(torch.int64),
|
||||
output_dtype=torch.bfloat16,
|
||||
quant_scales=[
|
||||
layer.w13_weight_scale.view(torch.int32),
|
||||
layer.mxfp4_weight_global_scale,
|
||||
layer.w2_weight_scale.view(torch.int32),
|
||||
layer.mxfp4_weight_global_scale,
|
||||
],
|
||||
input_sf=x_scale,
|
||||
fc1_expert_biases=layer.w13_weight_bias,
|
||||
fc2_expert_biases=layer.w2_weight_bias,
|
||||
swiglu_alpha=layer.swiglu_alpha,
|
||||
swiglu_beta=layer.swiglu_beta,
|
||||
swiglu_limit=layer.swiglu_limit,
|
||||
use_w4_group_scaling=False,
|
||||
use_mxfp8_act_scaling=True,
|
||||
activation_type=ActivationType.Swiglu,
|
||||
tune_max_num_tokens=8,
|
||||
output=expected,
|
||||
)
|
||||
assert torch.equal(actual, expected[:, :hidden].contiguous())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(pytest.main([__file__, "-v"]))
|
||||
|
||||
Reference in New Issue
Block a user