Migrate CompressedTensorsW4A4Nvfp4MoE TRT-LLM path onto MoeRunner (#32248)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
cff20a2fbb
commit
b83041c3cc
+29
-86
@@ -5,13 +5,8 @@ from typing import TYPE_CHECKING
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.distributed import get_tp_group
|
|
||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
|
||||||
use_symmetric_memory,
|
|
||||||
)
|
|
||||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
|
||||||
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
|
from sglang.srt.layers.moe import MoeRunner, MoeRunnerBackend, MoeRunnerConfig
|
||||||
from sglang.srt.layers.moe.utils import RoutingMethodType, get_moe_runner_backend
|
from sglang.srt.layers.moe.utils import get_moe_runner_backend
|
||||||
from sglang.srt.layers.quantization.compressed_tensors.schemes import (
|
from sglang.srt.layers.quantization.compressed_tensors.schemes import (
|
||||||
CompressedTensorsMoEScheme,
|
CompressedTensorsMoEScheme,
|
||||||
)
|
)
|
||||||
@@ -22,7 +17,7 @@ from sglang.srt.layers.quantization.utils import (
|
|||||||
replace_parameter,
|
replace_parameter,
|
||||||
swizzle_blockscale,
|
swizzle_blockscale,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import next_power_of_2, set_weight_attrs
|
from sglang.srt.utils import set_weight_attrs
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -282,7 +277,11 @@ class CompressedTensorsW4A4Nvfp4MoE(CompressedTensorsMoEScheme):
|
|||||||
):
|
):
|
||||||
self.moe_runner_config = moe_runner_config
|
self.moe_runner_config = moe_runner_config
|
||||||
if self.use_flashinfer_trtllm:
|
if self.use_flashinfer_trtllm:
|
||||||
self.runner = MoeRunner(MoeRunnerBackend.TRITON, moe_runner_config)
|
import sglang.srt.layers.moe.moe_runner.flashinfer_trtllm # noqa: F401 – triggers @register_fused_func
|
||||||
|
|
||||||
|
self.runner = MoeRunner(
|
||||||
|
MoeRunnerBackend.FLASHINFER_TRTLLM, moe_runner_config
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401 – triggers @register_fused_func
|
import sglang.srt.layers.moe.moe_runner.flashinfer_cutlass # noqa: F401 – triggers @register_fused_func
|
||||||
|
|
||||||
@@ -296,92 +295,38 @@ class CompressedTensorsW4A4Nvfp4MoE(CompressedTensorsMoEScheme):
|
|||||||
dispatch_output: StandardDispatchOutput,
|
dispatch_output: StandardDispatchOutput,
|
||||||
) -> CombineInput:
|
) -> CombineInput:
|
||||||
|
|
||||||
from sglang.srt.layers.moe.token_dispatcher import StandardCombineInput
|
|
||||||
|
|
||||||
x = dispatch_output.hidden_states
|
x = dispatch_output.hidden_states
|
||||||
topk_output = dispatch_output.topk_output
|
|
||||||
|
|
||||||
if self.use_flashinfer_trtllm:
|
if self.use_flashinfer_trtllm:
|
||||||
from flashinfer import trtllm_fp4_block_scale_moe
|
from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
|
||||||
|
FlashInferTrtllmFp4MoeQuantInfo,
|
||||||
from sglang.srt.layers.quantization.fp4_utils import fp4_quantize
|
|
||||||
|
|
||||||
router_logits = topk_output.router_logits
|
|
||||||
topk_config = topk_output.topk_config
|
|
||||||
|
|
||||||
# global_scale must be shape [1] (strict in cute-dsl backend).
|
|
||||||
hs_fp4_bytes, hs_sf_bytes = fp4_quantize(
|
|
||||||
x,
|
|
||||||
layer.w13_input_scale_quant[:1],
|
|
||||||
self.group_size, # sf_vec_size
|
|
||||||
False, # use_ue8m0
|
|
||||||
False, # is_sf_swizzled_layout
|
|
||||||
)
|
|
||||||
hs_fp4 = hs_fp4_bytes.reshape(x.shape[0], x.shape[1] // 2)
|
|
||||||
hs_scale = hs_sf_bytes.view(torch.float8_e4m3fn).reshape(
|
|
||||||
*hs_sf_bytes.shape[:-1], -1
|
|
||||||
)
|
|
||||||
|
|
||||||
correction_bias = (
|
|
||||||
None
|
|
||||||
if topk_config.correction_bias is None
|
|
||||||
else topk_config.correction_bias.to(x.dtype)
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert layer.routing_method_type is not None
|
assert layer.routing_method_type is not None
|
||||||
|
|
||||||
# DeepSeekV3 style routing requires float32 router logits
|
quant_info = FlashInferTrtllmFp4MoeQuantInfo(
|
||||||
if layer.routing_method_type == RoutingMethodType.DeepSeekV3:
|
w13_weight=layer.w13_weight,
|
||||||
router_logits = router_logits.to(torch.float32)
|
w2_weight=layer.w2_weight,
|
||||||
|
w13_weight_scale=layer.w13_weight_scale,
|
||||||
routed_scaling_factor = self.moe_runner_config.routed_scaling_factor
|
w2_weight_scale=layer.w2_weight_scale,
|
||||||
routed_scaling_factor = (
|
g1_scale_c=layer.g1_scale_c,
|
||||||
routed_scaling_factor if routed_scaling_factor is not None else 1.0
|
g1_alphas=layer.g1_alphas,
|
||||||
)
|
g2_alphas=layer.g2_alphas,
|
||||||
|
# global_scale must be shape [1]: the cute-dsl fp4_quantize
|
||||||
with use_symmetric_memory(
|
# backend reshapes it to [1]. process_weights_after_loading
|
||||||
get_tp_group(), disabled=not is_allocation_symmetric()
|
# expands this to [num_local_experts] (all equal), so slice the
|
||||||
):
|
# single value out (matches ModelOptNvFp4FusedMoEMethod, which
|
||||||
num_tokens = hs_fp4.shape[0]
|
# feeds a scalar).
|
||||||
hidden_size = (
|
w13_input_scale_quant=layer.w13_input_scale_quant[:1],
|
||||||
hs_fp4.shape[-1] * 2
|
global_num_experts=layer.num_experts,
|
||||||
if hs_fp4.dtype == torch.uint8
|
|
||||||
else hs_fp4.shape[-1]
|
|
||||||
)
|
|
||||||
symm_output = torch.empty(
|
|
||||||
num_tokens, hidden_size, dtype=torch.bfloat16, device=hs_fp4.device
|
|
||||||
)
|
|
||||||
|
|
||||||
output = trtllm_fp4_block_scale_moe(
|
|
||||||
routing_logits=router_logits,
|
|
||||||
routing_bias=correction_bias,
|
|
||||||
hidden_states=hs_fp4,
|
|
||||||
hidden_states_scale=hs_scale,
|
|
||||||
gemm1_weights=layer.w13_weight,
|
|
||||||
gemm1_weights_scale=layer.w13_weight_scale.view(torch.float8_e4m3fn),
|
|
||||||
gemm1_bias=None,
|
|
||||||
gemm1_alpha=None,
|
|
||||||
gemm1_beta=None,
|
|
||||||
gemm1_clamp_limit=None,
|
|
||||||
gemm2_weights=layer.w2_weight,
|
|
||||||
gemm2_weights_scale=layer.w2_weight_scale.view(torch.float8_e4m3fn),
|
|
||||||
gemm2_bias=None,
|
|
||||||
output1_scale_scalar=layer.g1_scale_c,
|
|
||||||
output1_scale_gate_scalar=layer.g1_alphas,
|
|
||||||
output2_scale_scalar=layer.g2_alphas,
|
|
||||||
num_experts=layer.num_experts,
|
|
||||||
top_k=topk_config.top_k,
|
|
||||||
n_group=topk_config.num_expert_group,
|
|
||||||
topk_group=topk_config.topk_group,
|
|
||||||
intermediate_size=layer.intermediate_size_per_partition,
|
|
||||||
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
|
local_expert_offset=layer.moe_ep_rank * layer.num_local_experts,
|
||||||
local_num_experts=layer.num_local_experts,
|
local_num_experts=layer.num_local_experts,
|
||||||
routed_scaling_factor=routed_scaling_factor,
|
intermediate_size_per_partition=layer.intermediate_size_per_partition,
|
||||||
routing_method_type=layer.routing_method_type,
|
routing_method_type=layer.routing_method_type,
|
||||||
do_finalize=True,
|
use_per_token_activation=False,
|
||||||
tune_max_num_tokens=next_power_of_2(hs_fp4.shape[0]),
|
gemm1_clamp_limit=None,
|
||||||
output=symm_output,
|
)
|
||||||
)[0]
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
else:
|
else:
|
||||||
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
from sglang.srt.layers.moe.moe_runner.flashinfer_cutlass import (
|
||||||
FlashInferCutlassMoeQuantInfo,
|
FlashInferCutlassMoeQuantInfo,
|
||||||
@@ -411,5 +356,3 @@ class CompressedTensorsW4A4Nvfp4MoE(CompressedTensorsMoEScheme):
|
|||||||
apply_routed_scaling_factor=False,
|
apply_routed_scaling_factor=False,
|
||||||
)
|
)
|
||||||
return self.runner.run(dispatch_output, quant_info)
|
return self.runner.run(dispatch_output, quant_info)
|
||||||
|
|
||||||
return StandardCombineInput(hidden_states=output)
|
|
||||||
|
|||||||
Reference in New Issue
Block a user