[DeepSeek V3] Reland: run routed experts on main stream in dual-stream MoE (#29463)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Khoa Pham
2026-06-29 15:05:01 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent 45314a9fcb
commit 6bdecb8206
+48 -45
View File
@@ -926,11 +926,13 @@ class DeepseekV2MoE(nn.Module):
*, *,
use_flashinfer_trtllm_bypass: bool = False, use_flashinfer_trtllm_bypass: bool = False,
) -> torch.Tensor: ) -> torch.Tensor:
# Note(kpham-sgl): launch the shared expert BEFORE the routed call.
# The routed deep_gemm pre-permute calls `dispose_tensor` which
# `set_()`s `hidden_states` to empty (host-side); any later kernel
# launch consuming `hidden_states` then captures `data_ptr() == 0`
# into the decode CUDA graph and replays from null.
current_stream = torch.cuda.current_stream() current_stream = torch.cuda.current_stream()
self.alt_stream.wait_stream(current_stream) self.alt_stream.wait_stream(current_stream)
shared_output = self._forward_shared_experts(
hidden_states, gemm_output_zero_allocator
)
server_args = get_global_server_args() server_args = get_global_server_args()
dispatch_info = ( dispatch_info = (
ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id) ExpertLocationDispatchInfo.init_new(layer_id=self.layer_id)
@@ -938,49 +940,50 @@ class DeepseekV2MoE(nn.Module):
else None else None
) )
with torch.cuda.stream(self.alt_stream): with torch.cuda.stream(self.alt_stream):
# router_logits: (num_tokens, n_experts) shared_output = self._forward_shared_experts(
router_logits = self.gate(hidden_states, gemm_output_zero_allocator) hidden_states, gemm_output_zero_allocator
if use_flashinfer_trtllm_bypass:
topk_output = BypassedTopKOutput(
hidden_states=hidden_states,
router_logits=router_logits,
topk_config=self.topk.topk_config,
)
else:
topk_kwargs = (
{"input_ids": input_ids_global}
if getattr(self, "is_hash", False)
else {}
)
topk_output = self.topk(
hidden_states,
router_logits,
expert_location_dispatch_info=dispatch_info,
**topk_kwargs,
)
deferred_finalize = (
shared_output is not None
and not self._shared_expert_tp1
and topk_output.format == TopKOutputFormat.BYPASSED
and self.experts.supports_deferred_finalize
) )
if deferred_finalize: # router_logits: (num_tokens, n_experts)
final_hidden_states = self.experts.forward_deferred_finalize( router_logits = self.gate(hidden_states, gemm_output_zero_allocator)
hidden_states, topk_output if use_flashinfer_trtllm_bypass:
) topk_output = BypassedTopKOutput(
elif use_flashinfer_trtllm_bypass: hidden_states=hidden_states,
final_hidden_states = self.experts.forward_impl( router_logits=router_logits,
hidden_states, topk_output topk_config=self.topk.topk_config,
) )
else: else:
final_hidden_states = self.experts(hidden_states, topk_output) topk_kwargs = (
if ( {"input_ids": input_ids_global}
not _is_cuda if getattr(self, "is_hash", False)
and not _is_musa else {}
and not _use_aiter )
or isinstance(self.experts.quant_method, KTEPWrapperMethod) topk_output = self.topk(
): hidden_states,
final_hidden_states *= self.routed_scaling_factor router_logits,
expert_location_dispatch_info=dispatch_info,
**topk_kwargs,
)
deferred_finalize = (
shared_output is not None
and not self._shared_expert_tp1
and topk_output.format == TopKOutputFormat.BYPASSED
and self.experts.supports_deferred_finalize
)
if deferred_finalize:
final_hidden_states = self.experts.forward_deferred_finalize(
hidden_states, topk_output
)
elif use_flashinfer_trtllm_bypass:
final_hidden_states = self.experts.forward_impl(hidden_states, topk_output)
else:
final_hidden_states = self.experts(hidden_states, topk_output)
if (
not _is_cuda
and not _is_musa
and not _use_aiter
or isinstance(self.experts.quant_method, KTEPWrapperMethod)
):
final_hidden_states *= self.routed_scaling_factor
current_stream.wait_stream(self.alt_stream) current_stream.wait_stream(self.alt_stream)