[AMD] Perf Kimi-K3 MoE optimization (#33838)
Co-authored-by: wunhuang <wunhuang@amd.com>
This commit is contained in:
@@ -16,7 +16,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import math
|
import math
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass, field
|
||||||
from enum import IntEnum, auto
|
from enum import IntEnum, auto
|
||||||
from typing import (
|
from typing import (
|
||||||
TYPE_CHECKING,
|
TYPE_CHECKING,
|
||||||
@@ -238,6 +238,32 @@ class TopKConfig:
|
|||||||
# Draft-side MoE blocks set this False so they never write the target's
|
# Draft-side MoE blocks set this False so they never write the target's
|
||||||
# process-global routed-experts capture buffer.
|
# process-global routed-experts capture buffer.
|
||||||
allow_routed_experts_capture: bool = True
|
allow_routed_experts_capture: bool = True
|
||||||
|
_correction_bias_dtype_cache: Optional[torch.Tensor] = field(
|
||||||
|
default=None, init=False, repr=False, compare=False
|
||||||
|
)
|
||||||
|
_correction_bias_cache_key: Optional[Tuple] = field(
|
||||||
|
default=None, init=False, repr=False, compare=False
|
||||||
|
)
|
||||||
|
|
||||||
|
def correction_bias_for_dtype(self, dtype: torch.dtype) -> Optional[torch.Tensor]:
|
||||||
|
"""Return correction bias in ``dtype``, reusing a per-TopK lazy copy."""
|
||||||
|
correction_bias = self.correction_bias
|
||||||
|
if correction_bias is None or correction_bias.dtype == dtype:
|
||||||
|
return correction_bias
|
||||||
|
|
||||||
|
# Weight loaders update parameters in place. Including the version in
|
||||||
|
# the key prevents an early access from retaining pre-load contents.
|
||||||
|
cache_key = (
|
||||||
|
correction_bias.data_ptr(),
|
||||||
|
correction_bias._version,
|
||||||
|
correction_bias.device,
|
||||||
|
correction_bias.dtype,
|
||||||
|
dtype,
|
||||||
|
)
|
||||||
|
if self._correction_bias_cache_key != cache_key:
|
||||||
|
self._correction_bias_dtype_cache = correction_bias.to(dtype=dtype)
|
||||||
|
self._correction_bias_cache_key = cache_key
|
||||||
|
return self._correction_bias_dtype_cache
|
||||||
|
|
||||||
|
|
||||||
# -------------------------------- TopKOutput ---------------------------------------
|
# -------------------------------- TopKOutput ---------------------------------------
|
||||||
@@ -2306,6 +2332,9 @@ def select_experts(
|
|||||||
info=expert_location_dispatch_info,
|
info=expert_location_dispatch_info,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
if _use_aiter and use_grouped_topk and correction_bias is not None:
|
||||||
|
correction_bias = topk_config.correction_bias_for_dtype(router_logits.dtype)
|
||||||
|
|
||||||
# DeepSeek V2/V3/R1 series models use grouped_top_k
|
# DeepSeek V2/V3/R1 series models use grouped_top_k
|
||||||
# remove num_fused_shared_experts from grouped_topk/biased_grouped_topk
|
# remove num_fused_shared_experts from grouped_topk/biased_grouped_topk
|
||||||
num_routed_topk = top_k - num_fused_shared_experts
|
num_routed_topk = top_k - num_fused_shared_experts
|
||||||
|
|||||||
@@ -156,6 +156,13 @@ _flashinfer_mxfp4_permute_indices_device_cache: dict[
|
|||||||
] = {}
|
] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def _aiter_situ_uses_gu_interleaved_weights() -> bool:
|
||||||
|
"""Match AITER's SiTU activation-mode precedence when choosing weight layout."""
|
||||||
|
a8w4 = get_bool_env_var("AITER_SITUV2_A8W4", "false")
|
||||||
|
a4w4 = get_bool_env_var("AITER_SITUV2_A4W4", "false")
|
||||||
|
return a8w4 or not a4w4
|
||||||
|
|
||||||
|
|
||||||
def _get_flashinfer_mxfp4_device_permute_indices(
|
def _get_flashinfer_mxfp4_device_permute_indices(
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
epilogue_tile_m: int,
|
epilogue_tile_m: int,
|
||||||
@@ -951,12 +958,17 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
.view(-1, n)
|
.view(-1, n)
|
||||||
)
|
)
|
||||||
|
|
||||||
k3_situ_a8w4 = (
|
# AITER selects the activation dtype at runtime. A8W4 takes precedence
|
||||||
os.environ.get("AITER_SITUV2_A8W4", "0") == "1"
|
# and, together with A16W4, uses the preshuffled GU-interleaved layout.
|
||||||
and getattr(layer.moe_runner_config, "activation", None) == "situ"
|
# A4W4 uses the generic separated layout instead; feeding it the
|
||||||
)
|
# A16/A8 layout makes real-checkpoint MoE outputs nearly orthogonal.
|
||||||
use_aiter_gu_interleave = k3_situ_a8w4 or (
|
k3_situ = getattr(layer.moe_runner_config, "activation", None) == "situ"
|
||||||
envs.SGLANG_USE_AITER_MOE_GU_ITLV.get() and gate_up_interleaved
|
use_aiter_gu_interleave = (
|
||||||
|
k3_situ and _aiter_situ_uses_gu_interleaved_weights()
|
||||||
|
) or (
|
||||||
|
not k3_situ
|
||||||
|
and envs.SGLANG_USE_AITER_MOE_GU_ITLV.get()
|
||||||
|
and gate_up_interleaved
|
||||||
)
|
)
|
||||||
if use_aiter_gu_interleave:
|
if use_aiter_gu_interleave:
|
||||||
layer.w13_weight.data = shuffle_weight_a16w4(layer.w13_weight, 16, True)
|
layer.w13_weight.data = shuffle_weight_a16w4(layer.w13_weight, 16, True)
|
||||||
@@ -1733,9 +1745,14 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
expanded_idx_to_permuted_idx=expanded_idx,
|
expanded_idx_to_permuted_idx=expanded_idx,
|
||||||
top_k=packed_topk.shape[1],
|
top_k=packed_topk.shape[1],
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
result = result[0]
|
|
||||||
return StandardCombineInput(hidden_states=result)
|
return StandardCombineInput(hidden_states=result)
|
||||||
|
# The finalized kernel writes to its explicit output
|
||||||
|
# argument. Do not propagate the FFI return tensor: some
|
||||||
|
# SiTU runner versions return a distinct wrapper/allocation
|
||||||
|
# even though symm_output contains the published result.
|
||||||
|
# Returning the destination makes the pointer contract
|
||||||
|
# explicit for K3's zero-copy latent buffer.
|
||||||
|
return StandardCombineInput(hidden_states=symm_output)
|
||||||
|
|
||||||
# Bypassed topk: route from logits inside the op.
|
# Bypassed topk: route from logits inside the op.
|
||||||
correction_bias = topk_output.topk_config.correction_bias
|
correction_bias = topk_output.topk_config.correction_bias
|
||||||
|
|||||||
@@ -0,0 +1,72 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.moe.topk import TopKConfig
|
||||||
|
from sglang.test.ci.ci_register import (
|
||||||
|
register_amd_ci,
|
||||||
|
register_cpu_ci,
|
||||||
|
register_cuda_ci,
|
||||||
|
)
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
register_cuda_ci(est_time=5, stage="base-a", runner_config="1-gpu-small")
|
||||||
|
register_amd_ci(est_time=5, stage="stage-a", runner_config="1-gpu-small-amd")
|
||||||
|
|
||||||
|
|
||||||
|
class CorrectionBiasCacheTestMixin:
|
||||||
|
device: str
|
||||||
|
|
||||||
|
def test_lazy_cache_uses_loaded_value_and_reuses_pointer(self):
|
||||||
|
correction_bias = torch.empty(8, dtype=torch.bfloat16, device=self.device)
|
||||||
|
config = TopKConfig(top_k=2, correction_bias=correction_bias)
|
||||||
|
|
||||||
|
loaded_value = torch.linspace(-1, 1, 8, dtype=torch.float32).to(
|
||||||
|
device=self.device, dtype=torch.bfloat16
|
||||||
|
)
|
||||||
|
with torch.no_grad():
|
||||||
|
correction_bias.copy_(loaded_value)
|
||||||
|
|
||||||
|
converted = config.correction_bias_for_dtype(torch.float32)
|
||||||
|
converted_again = config.correction_bias_for_dtype(torch.float32)
|
||||||
|
|
||||||
|
torch.testing.assert_close(converted, loaded_value.float())
|
||||||
|
self.assertEqual(converted.data_ptr(), converted_again.data_ptr())
|
||||||
|
self.assertNotEqual(converted.data_ptr(), correction_bias.data_ptr())
|
||||||
|
|
||||||
|
def test_cache_refreshes_if_weight_is_reloaded(self):
|
||||||
|
correction_bias = torch.zeros(8, dtype=torch.bfloat16, device=self.device)
|
||||||
|
config = TopKConfig(top_k=2, correction_bias=correction_bias)
|
||||||
|
converted = config.correction_bias_for_dtype(torch.float32)
|
||||||
|
|
||||||
|
reloaded_value = torch.arange(8, dtype=torch.float32).to(
|
||||||
|
device=self.device, dtype=torch.bfloat16
|
||||||
|
)
|
||||||
|
with torch.no_grad():
|
||||||
|
correction_bias.copy_(reloaded_value)
|
||||||
|
converted_after_reload = config.correction_bias_for_dtype(torch.float32)
|
||||||
|
|
||||||
|
torch.testing.assert_close(converted_after_reload, reloaded_value.float())
|
||||||
|
self.assertNotEqual(converted.data_ptr(), converted_after_reload.data_ptr())
|
||||||
|
|
||||||
|
def test_matching_dtype_returns_original_tensor(self):
|
||||||
|
correction_bias = torch.randn(8, dtype=torch.float32, device=self.device)
|
||||||
|
config = TopKConfig(top_k=2, correction_bias=correction_bias)
|
||||||
|
|
||||||
|
result = config.correction_bias_for_dtype(torch.float32)
|
||||||
|
|
||||||
|
self.assertEqual(result.data_ptr(), correction_bias.data_ptr())
|
||||||
|
|
||||||
|
|
||||||
|
class TestCorrectionBiasCacheCPU(CorrectionBiasCacheTestMixin, CustomTestCase):
|
||||||
|
device = "cpu"
|
||||||
|
|
||||||
|
|
||||||
|
@unittest.skipUnless(torch.cuda.is_available(), "needs a GPU")
|
||||||
|
class TestCorrectionBiasCacheGPU(CorrectionBiasCacheTestMixin, CustomTestCase):
|
||||||
|
device = "cuda"
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -0,0 +1,112 @@
|
|||||||
|
import sys
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers import zero_copy_context
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.standard import StandardDispatchOutput
|
||||||
|
from sglang.srt.layers.moe.topk import StandardTopKOutput
|
||||||
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
|
|
||||||
|
register_amd_ci(est_time=5, stage="stage-b", runner_config="1-gpu-small-amd")
|
||||||
|
|
||||||
|
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a GPU")
|
||||||
|
|
||||||
|
|
||||||
|
def test_situ_routed_moe_returns_published_output_buffer():
|
||||||
|
from sglang.srt.layers.moe import route_quant_handoff
|
||||||
|
from sglang.srt.layers.quantization import mxfp4 as mxfp4_module
|
||||||
|
from sglang.srt.layers.quantization.mxfp4 import Mxfp4MoEMethod
|
||||||
|
|
||||||
|
tokens, hidden, top_k = 3, 128, 2
|
||||||
|
x = torch.randn(tokens, hidden, dtype=torch.bfloat16, device="cuda")
|
||||||
|
x_quant = torch.zeros(tokens, hidden, dtype=torch.uint8, device="cuda")
|
||||||
|
x_scale = torch.zeros(tokens, hidden // 32, dtype=torch.uint8, device="cuda")
|
||||||
|
packed_topk = torch.zeros(tokens, top_k, dtype=torch.int32, device="cuda")
|
||||||
|
topk_output = StandardTopKOutput(
|
||||||
|
topk_weights=torch.full(
|
||||||
|
(tokens, top_k), 0.5, dtype=torch.float32, device="cuda"
|
||||||
|
),
|
||||||
|
topk_ids=torch.zeros(tokens, top_k, dtype=torch.int32, device="cuda"),
|
||||||
|
router_logits=torch.empty(tokens, 0, dtype=torch.float32, device="cuda"),
|
||||||
|
)
|
||||||
|
dispatch_output = StandardDispatchOutput(
|
||||||
|
hidden_states=x,
|
||||||
|
hidden_states_scale=None,
|
||||||
|
topk_output=topk_output,
|
||||||
|
)
|
||||||
|
|
||||||
|
method = Mxfp4MoEMethod.__new__(Mxfp4MoEMethod)
|
||||||
|
method.use_deep_gemm = False
|
||||||
|
method.use_marlin = False
|
||||||
|
method.use_flashinfer = True
|
||||||
|
method._fi_kernel = None
|
||||||
|
method.flashinfer_mxfp4_moe_precision = "default"
|
||||||
|
method.hidden_size = hidden
|
||||||
|
method.intermediate_size_per_partition = 128
|
||||||
|
method.moe_runner_config = SimpleNamespace(activation="situ")
|
||||||
|
|
||||||
|
dummy = torch.empty(1, dtype=torch.uint8, device="cuda")
|
||||||
|
layer = SimpleNamespace(
|
||||||
|
moe_ep_rank=0,
|
||||||
|
num_local_experts=1,
|
||||||
|
num_experts=1,
|
||||||
|
w13_weight=dummy,
|
||||||
|
w13_weight_scale=dummy,
|
||||||
|
gemm1_alpha=None,
|
||||||
|
gemm1_clamp_limit=None,
|
||||||
|
w2_weight=dummy,
|
||||||
|
w2_weight_scale=dummy,
|
||||||
|
)
|
||||||
|
expected = (
|
||||||
|
torch.arange(tokens * hidden, dtype=torch.float32, device="cuda")
|
||||||
|
.reshape(tokens, hidden)
|
||||||
|
.to(torch.bfloat16)
|
||||||
|
)
|
||||||
|
returned_ptr = None
|
||||||
|
|
||||||
|
def fake_routed_moe(**kwargs):
|
||||||
|
nonlocal returned_ptr
|
||||||
|
kwargs["output"].copy_(expected)
|
||||||
|
ffi_result = kwargs["output"].clone()
|
||||||
|
returned_ptr = ffi_result.data_ptr()
|
||||||
|
return ffi_result
|
||||||
|
|
||||||
|
latent = torch.empty_like(x)
|
||||||
|
with (
|
||||||
|
patch.object(
|
||||||
|
route_quant_handoff,
|
||||||
|
"take",
|
||||||
|
return_value=(packed_topk, x_quant, x_scale),
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"sglang.srt.layers.quantization.mxfp4.trtllm_fp4_block_scale_routed_moe",
|
||||||
|
side_effect=fake_routed_moe,
|
||||||
|
create=True,
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
mxfp4_module,
|
||||||
|
"RoutingMethodType",
|
||||||
|
SimpleNamespace(TopK=SimpleNamespace(value=0)),
|
||||||
|
create=True,
|
||||||
|
),
|
||||||
|
patch.object(
|
||||||
|
mxfp4_module,
|
||||||
|
"ActivationType",
|
||||||
|
SimpleNamespace(Situ=SimpleNamespace(value=0)),
|
||||||
|
create=True,
|
||||||
|
),
|
||||||
|
zero_copy_context.set_moe_output(latent),
|
||||||
|
):
|
||||||
|
combine_input = method.apply(layer, dispatch_output)
|
||||||
|
|
||||||
|
assert returned_ptr is not None
|
||||||
|
assert returned_ptr != latent.data_ptr()
|
||||||
|
assert combine_input.hidden_states.data_ptr() == latent.data_ptr()
|
||||||
|
torch.testing.assert_close(combine_input.hidden_states, expected, rtol=0, atol=0)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from sglang.test.ci.ci_register import register_amd_ci
|
||||||
|
|
||||||
|
register_amd_ci(est_time=5, stage="stage-b", runner_config="1-gpu-small-amd")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("a8w4", "a4w4", "expected"),
|
||||||
|
[
|
||||||
|
(None, None, True),
|
||||||
|
("0", "1", False),
|
||||||
|
("1", "0", True),
|
||||||
|
("1", "1", True),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_aiter_situ_weight_layout_matches_activation_mode_precedence(
|
||||||
|
monkeypatch, a8w4, a4w4, expected
|
||||||
|
):
|
||||||
|
from sglang.srt.layers.quantization.mxfp4 import (
|
||||||
|
_aiter_situ_uses_gu_interleaved_weights,
|
||||||
|
)
|
||||||
|
|
||||||
|
monkeypatch.delenv("AITER_SITUV2_A8W4", raising=False)
|
||||||
|
monkeypatch.delenv("AITER_SITUV2_A4W4", raising=False)
|
||||||
|
if a8w4 is not None:
|
||||||
|
monkeypatch.setenv("AITER_SITUV2_A8W4", a8w4)
|
||||||
|
if a4w4 is not None:
|
||||||
|
monkeypatch.setenv("AITER_SITUV2_A4W4", a4w4)
|
||||||
|
|
||||||
|
assert _aiter_situ_uses_gu_interleaved_weights() is expected
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
Reference in New Issue
Block a user