When attention TP for linear and full attention, use Flashinfer allreduce fusion (#29699)
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
co-authored by
Brayden Zhong
parent
52c6e27e7e
commit
8f40b5eb3f
@@ -742,6 +742,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
mup_vector: Optional[torch.Tensor] = None,
|
mup_vector: Optional[torch.Tensor] = None,
|
||||||
use_triton_causal_conv: bool = False,
|
use_triton_causal_conv: bool = False,
|
||||||
|
should_allreduce_fusion: bool = False,
|
||||||
):
|
):
|
||||||
assert isinstance(self.forward_metadata, Mamba2Metadata)
|
assert isinstance(self.forward_metadata, Mamba2Metadata)
|
||||||
# Page-major stores state strided; only the stride-aware Triton causal-conv
|
# Page-major stores state strided; only the stride-aware Triton causal-conv
|
||||||
@@ -759,6 +760,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
|||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
mup_vector=mup_vector,
|
mup_vector=mup_vector,
|
||||||
use_triton_causal_conv=use_triton_causal_conv,
|
use_triton_causal_conv=use_triton_causal_conv,
|
||||||
|
should_allreduce_fusion=should_allreduce_fusion,
|
||||||
)
|
)
|
||||||
|
|
||||||
if forward_batch.mamba_track_mask is not None:
|
if forward_batch.mamba_track_mask is not None:
|
||||||
|
|||||||
@@ -448,6 +448,7 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
forward_batch: ForwardBatch,
|
forward_batch: ForwardBatch,
|
||||||
mup_vector: Optional[torch.Tensor] = None,
|
mup_vector: Optional[torch.Tensor] = None,
|
||||||
use_triton_causal_conv: bool = False,
|
use_triton_causal_conv: bool = False,
|
||||||
|
should_allreduce_fusion: bool = False,
|
||||||
):
|
):
|
||||||
# Returns the projected result. When `output` is given it is also
|
# Returns the projected result. When `output` is given it is also
|
||||||
# written into that buffer (required by the cuda-graph split ops, which
|
# written into that buffer (required by the cuda-graph split ops, which
|
||||||
@@ -760,7 +761,9 @@ class MambaMixer2(torch.nn.Module):
|
|||||||
# norm usage
|
# norm usage
|
||||||
hidden_states = self.norm(preallocated_ssm_out, gate)
|
hidden_states = self.norm(preallocated_ssm_out, gate)
|
||||||
|
|
||||||
mixer_out, _ = self.out_proj(hidden_states)
|
mixer_out, _ = self.out_proj(
|
||||||
|
hidden_states, skip_all_reduce=should_allreduce_fusion
|
||||||
|
)
|
||||||
if output is not None:
|
if output is not None:
|
||||||
output[:padded_num_tokens].copy_(mixer_out)
|
output[:padded_num_tokens].copy_(mixer_out)
|
||||||
|
|
||||||
|
|||||||
@@ -495,11 +495,18 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
self.layer_communicator = make_layer_communicator(self.norm, for_attn=True)
|
self.layer_communicator = make_layer_communicator(
|
||||||
|
self.norm,
|
||||||
|
for_attn=True,
|
||||||
|
is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1,
|
||||||
|
)
|
||||||
self._set_prev_layer_is_attn(config, layer_idx)
|
self._set_prev_layer_is_attn(config, layer_idx)
|
||||||
|
|
||||||
def _forward_mamba(
|
def _forward_mamba(
|
||||||
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
should_allreduce_fusion: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
"""Core Mamba forward logic, called directly or via split op."""
|
"""Core Mamba forward logic, called directly or via split op."""
|
||||||
original_num_tokens = hidden_states.shape[0]
|
original_num_tokens = hidden_states.shape[0]
|
||||||
@@ -517,6 +524,7 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
|||||||
output=None,
|
output=None,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
use_triton_causal_conv=True,
|
use_triton_causal_conv=True,
|
||||||
|
should_allreduce_fusion=should_allreduce_fusion,
|
||||||
)
|
)
|
||||||
return pad_to_original_num_tokens(output, original_num_tokens)
|
return pad_to_original_num_tokens(output, original_num_tokens)
|
||||||
|
|
||||||
@@ -544,17 +552,29 @@ class NemotronHMambaDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
|||||||
self.norm, hidden_states, residual
|
self.norm, hidden_states, residual
|
||||||
)
|
)
|
||||||
|
|
||||||
|
should_allreduce_fusion = (
|
||||||
|
self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer(
|
||||||
|
forward_batch
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
if is_in_breakable_cuda_graph():
|
if is_in_breakable_cuda_graph():
|
||||||
output = torch.empty_like(hidden_states)
|
output = torch.empty_like(hidden_states)
|
||||||
breakable_nemotron_mamba2_with_output(hidden_states, output, self.layer_id)
|
breakable_nemotron_mamba2_with_output(
|
||||||
return output, residual
|
hidden_states, output, self.layer_id, should_allreduce_fusion
|
||||||
|
)
|
||||||
if is_in_tc_piecewise_cuda_graph():
|
elif is_in_tc_piecewise_cuda_graph():
|
||||||
output = torch.empty_like(hidden_states)
|
output = torch.empty_like(hidden_states)
|
||||||
nemotron_mamba2_with_output(hidden_states, output, self.layer_id)
|
nemotron_mamba2_with_output(
|
||||||
return output, residual
|
hidden_states, output, self.layer_id, should_allreduce_fusion
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
output = self._forward_mamba(hidden_states, forward_batch)
|
output = self._forward_mamba(
|
||||||
|
hidden_states, forward_batch, should_allreduce_fusion
|
||||||
|
)
|
||||||
|
|
||||||
|
if should_allreduce_fusion:
|
||||||
|
output._sglang_needs_allreduce_fusion = True
|
||||||
return output, residual
|
return output, residual
|
||||||
|
|
||||||
|
|
||||||
@@ -625,13 +645,18 @@ class NemotronHAttention(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
|
self,
|
||||||
|
hidden_states: torch.Tensor,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
should_allreduce_fusion: bool = False,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
if not is_dp_attention_enabled():
|
if not is_dp_attention_enabled():
|
||||||
qkv, _ = self.qkv_proj(hidden_states)
|
qkv, _ = self.qkv_proj(hidden_states)
|
||||||
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
|
||||||
attn_output = self.attn.forward(q, k, v, forward_batch)
|
attn_output = self.attn.forward(q, k, v, forward_batch)
|
||||||
output, _ = self.o_proj(attn_output)
|
output, _ = self.o_proj(
|
||||||
|
attn_output, skip_all_reduce=should_allreduce_fusion
|
||||||
|
)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
padded_shape = hidden_states.shape[0]
|
padded_shape = hidden_states.shape[0]
|
||||||
@@ -682,7 +707,11 @@ class NemotronHAttentionDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
self.norm = RMSNorm(config.hidden_size, eps=config.layer_norm_epsilon)
|
||||||
self.layer_communicator = make_layer_communicator(self.norm, for_attn=True)
|
self.layer_communicator = make_layer_communicator(
|
||||||
|
self.norm,
|
||||||
|
for_attn=True,
|
||||||
|
is_last_layer=layer_idx == len(config.hybrid_override_pattern) - 1,
|
||||||
|
)
|
||||||
self._set_prev_layer_is_attn(config, layer_idx)
|
self._set_prev_layer_is_attn(config, layer_idx)
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
@@ -705,9 +734,19 @@ class NemotronHAttentionDecoderLayer(NemotronHAttnLikeDecoderLayer):
|
|||||||
self.norm, hidden_states, residual
|
self.norm, hidden_states, residual
|
||||||
)
|
)
|
||||||
|
|
||||||
hidden_states = self.mixer.forward(
|
should_allreduce_fusion = (
|
||||||
hidden_states=hidden_states, forward_batch=forward_batch
|
self.layer_communicator.should_fuse_mlp_allreduce_with_next_layer(
|
||||||
|
forward_batch
|
||||||
)
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
hidden_states = self.mixer.forward(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
forward_batch=forward_batch,
|
||||||
|
should_allreduce_fusion=should_allreduce_fusion,
|
||||||
|
)
|
||||||
|
if should_allreduce_fusion:
|
||||||
|
hidden_states._sglang_needs_allreduce_fusion = True
|
||||||
return hidden_states, residual
|
return hidden_states, residual
|
||||||
|
|
||||||
|
|
||||||
@@ -1168,6 +1207,7 @@ def nemotron_mamba2_with_output(
|
|||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
output: torch.Tensor,
|
output: torch.Tensor,
|
||||||
layer_id: int,
|
layer_id: int,
|
||||||
|
should_allreduce_fusion: bool = False,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Split op for Mamba2 forward in piecewise CUDA graph mode."""
|
"""Split op for Mamba2 forward in piecewise CUDA graph mode."""
|
||||||
context = get_tc_piecewise_forward_context()
|
context = get_tc_piecewise_forward_context()
|
||||||
@@ -1187,7 +1227,9 @@ def nemotron_mamba2_with_output(
|
|||||||
if hidden_states.shape[0] != num_actual_tokens:
|
if hidden_states.shape[0] != num_actual_tokens:
|
||||||
hidden_states = hidden_states[:num_actual_tokens]
|
hidden_states = hidden_states[:num_actual_tokens]
|
||||||
|
|
||||||
ret = mamba_layer._forward_mamba(hidden_states, forward_batch)
|
ret = mamba_layer._forward_mamba(
|
||||||
|
hidden_states, forward_batch, should_allreduce_fusion
|
||||||
|
)
|
||||||
|
|
||||||
# Copy result back; output may be larger (padded) so only fill actual tokens
|
# Copy result back; output may be larger (padded) so only fill actual tokens
|
||||||
output[:num_actual_tokens].view(ret.shape).copy_(ret)
|
output[:num_actual_tokens].view(ret.shape).copy_(ret)
|
||||||
|
|||||||
Reference in New Issue
Block a user