Remove redundant cast and copy in calling trtllm_fp8_block_scale_moe (#28555)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-06-18 18:11:29 -07:00
committed by GitHub
co-authored by Brayden Zhong
parent ef01618dfb
commit 05ee93c44f
@@ -652,11 +652,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
if TopKOutputChecker.format_is_bypassed(topk_output): if TopKOutputChecker.format_is_bypassed(topk_output):
router_logits = topk_output.router_logits router_logits = topk_output.router_logits
topk_config = topk_output.topk_config topk_config = topk_output.topk_config
correction_bias = ( correction_bias = topk_config.correction_bias
None
if topk_config.correction_bias is None
else topk_config.correction_bias.to(hidden_states.dtype)
)
else: else:
router_logits = None router_logits = None
topk_config = None topk_config = None
@@ -685,9 +681,9 @@ def fused_experts_none_to_flashinfer_trtllm_fp8(
a_sf_t = a_sf.view(torch.uint8).reshape(hidden_states.shape[0], -1) a_sf_t = a_sf.view(torch.uint8).reshape(hidden_states.shape[0], -1)
else: else:
a_q, a_sf = per_token_group_quant_fp8( a_q, a_sf = per_token_group_quant_fp8(
hidden_states, quant_info.weight_block_k hidden_states, quant_info.weight_block_k, column_major_scales=True
) )
a_sf_t = a_sf.t().contiguous() a_sf_t = a_sf.t()
# Allocate output inside symmetric memory context # Allocate output inside symmetric memory context
with use_symmetric_memory( with use_symmetric_memory(
@@ -1053,11 +1049,7 @@ def fused_experts_none_to_flashinfer_trtllm_fp4(
topk_config = topk_output.topk_config topk_config = topk_output.topk_config
routing_method_type = quant_info.routing_method_type routing_method_type = quant_info.routing_method_type
correction_bias = ( correction_bias = topk_config.correction_bias
None
if topk_config.correction_bias is None
else topk_config.correction_bias.to(hidden_states.dtype)
)
moe_kwargs = dict( moe_kwargs = dict(
routing_logits=router_logits, routing_logits=router_logits,
routing_bias=correction_bias, routing_bias=correction_bias,