[Kimi K2.5] Fix eagle3 aux capture for tp>1 when AR fusion is enabled (#28343)
This commit is contained in:
@@ -498,11 +498,13 @@ class LayerCommunicator:
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
||||||
post_residual_addition: Optional[torch.Tensor] = None,
|
post_residual_addition: Optional[torch.Tensor] = None,
|
||||||
|
quant_format: str = "",
|
||||||
):
|
):
|
||||||
hidden_states, residual = self.prepare_attn(
|
hidden_states, residual = self.prepare_attn(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
residual,
|
residual,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
|
quant_format=quant_format,
|
||||||
post_residual_addition=post_residual_addition,
|
post_residual_addition=post_residual_addition,
|
||||||
)
|
)
|
||||||
if captured_last_layer_outputs is not None:
|
if captured_last_layer_outputs is not None:
|
||||||
@@ -511,12 +513,39 @@ class LayerCommunicator:
|
|||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
context=self._context,
|
context=self._context,
|
||||||
)
|
)
|
||||||
if gathered_last_layer_output is residual:
|
if (
|
||||||
# Clone to avoid modifying the original residual by Custom RMSNorm inplace operation
|
gathered_last_layer_output is residual
|
||||||
|
and not self._post_attn_residual_is_read_only(residual)
|
||||||
|
):
|
||||||
gathered_last_layer_output = residual.clone()
|
gathered_last_layer_output = residual.clone()
|
||||||
captured_last_layer_outputs.append(gathered_last_layer_output)
|
captured_last_layer_outputs.append(gathered_last_layer_output)
|
||||||
return hidden_states, residual
|
return hidden_states, residual
|
||||||
|
|
||||||
|
def _post_attn_residual_is_read_only(self, residual: torch.Tensor) -> bool:
|
||||||
|
"""True if ``prepare_mlp``'s post-attention RMSNorm leaves ``residual``
|
||||||
|
untouched, so Eagle3 aux capture can keep its reference and skip the clone.
|
||||||
|
|
||||||
|
Only the flashinfer all-reduce-fusion path writes a fresh ``residual_out``
|
||||||
|
(see ``flashinfer_allreduce_residual_rmsnorm``); the aiter fused kernel and
|
||||||
|
every plain norm fold into ``residual`` in place. That path is reachable
|
||||||
|
only from the ``_gather_*`` communicate-fns, and only when they fall past
|
||||||
|
their input-scattered branch.
|
||||||
|
"""
|
||||||
|
norm_fn = getattr(
|
||||||
|
self._communicate_with_all_reduce_and_layer_norm_fn,
|
||||||
|
"func",
|
||||||
|
self._communicate_with_all_reduce_and_layer_norm_fn,
|
||||||
|
)
|
||||||
|
uses_gather_norm = norm_fn in (
|
||||||
|
CommunicateWithAllReduceAndLayerNormFn._gather_hidden_states_and_residual,
|
||||||
|
CommunicateWithAllReduceAndLayerNormFn._gather_hidden_states_and_residual_moe,
|
||||||
|
)
|
||||||
|
return (
|
||||||
|
uses_gather_norm
|
||||||
|
and not get_attn_tp_context().input_scattered
|
||||||
|
and apply_flashinfer_allreduce_fusion(residual.shape[0])
|
||||||
|
)
|
||||||
|
|
||||||
def prepare_attn(
|
def prepare_attn(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
|
|||||||
@@ -75,7 +75,6 @@ from sglang.srt.layers.communicator_dsa_cp import DSACPLayerCommunicator
|
|||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
get_attention_cp_rank,
|
get_attention_cp_rank,
|
||||||
get_attention_cp_size,
|
get_attention_cp_size,
|
||||||
get_attention_tp_group,
|
|
||||||
get_attention_tp_rank,
|
get_attention_tp_rank,
|
||||||
get_attention_tp_size,
|
get_attention_tp_size,
|
||||||
)
|
)
|
||||||
@@ -2130,13 +2129,17 @@ class DeepseekV2DecoderLayer(nn.Module):
|
|||||||
gemm_output_zero_allocator: BumpAllocator = None,
|
gemm_output_zero_allocator: BumpAllocator = None,
|
||||||
llama_4_scaling: Optional[torch.Tensor] = None,
|
llama_4_scaling: Optional[torch.Tensor] = None,
|
||||||
prev_topk_indices: Optional[torch.Tensor] = None,
|
prev_topk_indices: Optional[torch.Tensor] = None,
|
||||||
|
captured_last_layer_outputs: Optional[List[torch.Tensor]] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
hidden_states_orig = hidden_states
|
hidden_states_orig = hidden_states
|
||||||
hidden_states, residual = self.layer_communicator.prepare_attn(
|
hidden_states, residual = (
|
||||||
|
self.layer_communicator.prepare_attn_and_capture_last_layer_outputs(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
residual,
|
residual,
|
||||||
forward_batch,
|
forward_batch,
|
||||||
getattr(self, "_gfx95_quant_format", ""),
|
captured_last_layer_outputs=captured_last_layer_outputs,
|
||||||
|
quant_format=getattr(self, "_gfx95_quant_format", ""),
|
||||||
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
hidden_states = self.self_attn(
|
hidden_states = self.self_attn(
|
||||||
@@ -2483,14 +2486,6 @@ class DeepseekV2Model(nn.Module):
|
|||||||
else get_global_expert_distribution_recorder().with_current_layer(i)
|
else get_global_expert_distribution_recorder().with_current_layer(i)
|
||||||
)
|
)
|
||||||
with ctx:
|
with ctx:
|
||||||
if i in self.layers_to_capture:
|
|
||||||
if self.enable_a2a_moe and i > self.first_k_dense_replace:
|
|
||||||
aux_hidden_state = get_attention_tp_group().all_gather(
|
|
||||||
hidden_states + residual, dim=0
|
|
||||||
)
|
|
||||||
aux_hidden_states.append(aux_hidden_state)
|
|
||||||
else:
|
|
||||||
aux_hidden_states.append(hidden_states + residual)
|
|
||||||
layer = self.layers[i]
|
layer = self.layers[i]
|
||||||
hidden_states, residual, topk_indices = layer(
|
hidden_states, residual, topk_indices = layer(
|
||||||
positions,
|
positions,
|
||||||
@@ -2501,6 +2496,9 @@ class DeepseekV2Model(nn.Module):
|
|||||||
gemm_output_zero_allocator,
|
gemm_output_zero_allocator,
|
||||||
llama_4_scaling,
|
llama_4_scaling,
|
||||||
prev_topk_indices=topk_indices,
|
prev_topk_indices=topk_indices,
|
||||||
|
captured_last_layer_outputs=(
|
||||||
|
aux_hidden_states if i in self.layers_to_capture else None
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
if normal_end_layer != self.end_layer:
|
if normal_end_layer != self.end_layer:
|
||||||
|
|||||||
Reference in New Issue
Block a user