From 388018a5bd41a25f55c209241525141111e27382 Mon Sep 17 00:00:00 2001 From: liupeng374 <782420244@qq.com> Date: Thu, 11 Dec 2025 20:12:59 +0800 Subject: [PATCH] [NPU] adapt dsv3.2 nsa prefill context parallel (#14541) --- .../sglang/srt/distributed/parallel_state.py | 9 +- .../npu/attention/ascend_backend.py | 135 +++++++++-- .../modules/deepseek_v2_attention_mla_npu.py | 9 + .../sglang/srt/hardware_backend/npu/utils.py | 8 + .../srt/layers/attention/nsa/nsa_indexer.py | 211 ++++++++++-------- .../sglang/srt/layers/attention/nsa/utils.py | 29 ++- .../sglang/srt/layers/communicator_nsa_cp.py | 15 +- test/srt/ascend/test_ascend_tp4_bf16.py | 1 + 8 files changed, 282 insertions(+), 135 deletions(-) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index cf90f6fe0..216b8d5a3 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -761,13 +761,10 @@ class GroupCoordinator: stream is not None ), f"Invalid params stream ({stream}, Please specify the stream to use when calling cp_all_gather_into_tensor_async.)" pynccl_comm = self.pynccl_comm - if pynccl_comm is not None: - pynccl_comm.cp_all_gather_into_tensor(output, input, stream=stream) + if pynccl_comm is None or pynccl_comm.disabled: + self.all_gather_into_tensor(output, input) else: - logger.warning("not all_gather_into_tensor_async") - torch.ops.sglang.reg_all_gather_into_tensor( - output, input, group_name=self.unique_name - ) + pynccl_comm.cp_all_gather_into_tensor(output, input, stream=stream) def all_gather( self, diff --git a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py index 2c804058a..0e21e680c 100644 --- a/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py +++ b/python/sglang/srt/hardware_backend/npu/attention/ascend_backend.py @@ -11,6 +11,7 @@ from sglang.srt.hardware_backend.npu.attention.mla_preprocess import ( is_mla_preprocess_enabled, ) from sglang.srt.layers.attention.base_attn_backend import AttentionBackend +from sglang.srt.layers.attention.nsa.utils import is_nsa_enable_prefill_cp from sglang.srt.layers.attention.torch_native_backend import TorchNativeAttnBackend from sglang.srt.layers.dp_attention import get_attention_tp_size from sglang.srt.layers.radix_attention import AttentionType @@ -42,6 +43,7 @@ class ForwardMetadata: seq_lens_list_cumsum: Optional[List[int]] = None seq_lens: Optional[torch.Tensor] = None actual_seq_lengths_q: Optional[torch.Tensor] = None + actual_seq_lengths_kv: Optional[torch.Tensor] = None # prefix cache prefix_lens: Optional[torch.Tensor] = None @@ -267,7 +269,6 @@ class AscendAttnBackend(AttentionBackend): def init_forward_metadata(self, forward_batch: ForwardBatch): """Init the metadata for a forward pass.""" - tp_size = get_attention_tp_size() self.forward_metadata = ForwardMetadata() seq_lens_max = forward_batch.seq_lens.max() if forward_batch.forward_mode.is_target_verify(): @@ -411,6 +412,72 @@ class AscendAttnBackend(AttentionBackend): def get_cuda_graph_seq_len_fill_value(self): return 0 + def do_cp_balance_attn( + self, + q_nope, + k_nope, + q_pe, + k_pe, + topk_indices, + layer, + actual_seq_qlen, + actual_seq_lengths_kv, + ): + seq_len = q_nope.shape[0] + split_len = (seq_len + 1) // 2 + q_nope_prev, q_nope_next = torch.split(q_nope, split_len, dim=0) + q_rope_prev, q_rope_next = torch.split(q_pe, split_len, dim=0) + q_nope_prev = q_nope_prev.contiguous() + q_nope_next = q_nope_next.contiguous() + q_rope_prev = q_rope_prev.contiguous() + q_rope_next = q_rope_next.contiguous() + topk_indices_prev, topk_indices_next = topk_indices + + actual_seq_qlen_prev, actual_seq_qlen_next = actual_seq_qlen + actual_seq_lengths_kv_prev, actual_seq_lengths_kv_next = actual_seq_lengths_kv + + attn_out_prev = torch.ops.custom.npu_sparse_flash_attention( + query=q_nope_prev, + key=k_nope, + value=k_nope, + query_rope=q_rope_prev, + key_rope=k_pe, + sparse_indices=topk_indices_prev, + scale_value=layer.scaling, + actual_seq_lengths_query=actual_seq_qlen_prev.to( + device=q_nope.device, dtype=torch.int32 + ), + actual_seq_lengths_kv=actual_seq_lengths_kv_prev.to( + device=q_nope.device, dtype=torch.int32 + ), + block_table=self.forward_metadata.block_tables, + sparse_block_size=1, + layout_query="TND", + layout_kv="PA_BSND", + sparse_mode=3, + ) + attn_out_next = torch.ops.custom.npu_sparse_flash_attention( + query=q_nope_next, + key=k_nope, + value=k_nope, + query_rope=q_rope_next, + key_rope=k_pe, + sparse_indices=topk_indices_next, + scale_value=layer.scaling, + actual_seq_lengths_query=actual_seq_qlen_next.to( + device=q_nope.device, dtype=torch.int32 + ), + actual_seq_lengths_kv=actual_seq_lengths_kv_next.to( + device=q_nope.device, dtype=torch.int32 + ), + block_table=self.forward_metadata.block_tables, + sparse_block_size=1, + layout_query="TND", + layout_kv="PA_BSND", + sparse_mode=3, + ) + return torch.cat([attn_out_prev, attn_out_next], dim=0) + def forward_sparse( self, q: torch.Tensor, @@ -440,9 +507,12 @@ class AscendAttnBackend(AttentionBackend): ) q_nope, q_pe = q, q_rope k_nope, k_pe = forward_batch.token_to_kv_pool.get_kv_buffer(layer.layer_id) - block_table = self.forward_metadata.block_tables + if is_prefill: - actual_seq_qlen = torch.cumsum(forward_batch.seq_lens, dim=0) + if self.forward_metadata.actual_seq_lengths_q is not None: + actual_seq_qlen = self.forward_metadata.actual_seq_lengths_q + else: + actual_seq_qlen = torch.cumsum(forward_batch.seq_lens, dim=0) else: if self.forward_metadata.actual_seq_lengths_q is None: if ( @@ -471,27 +541,46 @@ class AscendAttnBackend(AttentionBackend): ) else: actual_seq_qlen = self.forward_metadata.actual_seq_lengths_q - if self.forward_metadata.seq_lens_cpu_int is None: - actual_seq_lengths_kv = self.forward_metadata.seq_lens - else: - actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_int - attn_out = torch.ops.custom.npu_sparse_flash_attention( - query=q_nope, - key=k_nope, - value=k_nope, - query_rope=q_pe, - key_rope=k_pe, - sparse_indices=topk_indices, - scale_value=layer.scaling, - actual_seq_lengths_query=actual_seq_qlen.to(torch.int32), - actual_seq_lengths_kv=actual_seq_lengths_kv.to(q.device), - block_table=block_table, - sparse_block_size=1, - layout_query="TND", - layout_kv="PA_BSND", - sparse_mode=3, - ) + if self.forward_metadata.actual_seq_lengths_kv is not None: + actual_seq_lengths_kv = self.forward_metadata.actual_seq_lengths_kv + elif self.forward_metadata.seq_lens_cpu_int is not None: + actual_seq_lengths_kv = self.forward_metadata.seq_lens_cpu_int + else: + actual_seq_lengths_kv = self.forward_metadata.seq_lens + + if ( + is_prefill + and is_nsa_enable_prefill_cp() + and forward_batch.nsa_cp_metadata is not None + ): + attn_out = self.do_cp_balance_attn( + q_nope, + k_nope, + q_pe, + k_pe, + topk_indices, + layer, + actual_seq_qlen, + actual_seq_lengths_kv, + ) + else: + attn_out = torch.ops.custom.npu_sparse_flash_attention( + query=q_nope, + key=k_nope, + value=k_nope, + query_rope=q_pe, + key_rope=k_pe, + sparse_indices=topk_indices, + scale_value=layer.scaling, + actual_seq_lengths_query=actual_seq_qlen, + actual_seq_lengths_kv=actual_seq_lengths_kv.to(q.device), + block_table=self.forward_metadata.block_tables, + sparse_block_size=1, + layout_query="TND", + layout_kv="PA_BSND", + sparse_mode=3, + ) return attn_out diff --git a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py index 470617585..8b916d8de 100644 --- a/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py +++ b/python/sglang/srt/hardware_backend/npu/modules/deepseek_v2_attention_mla_npu.py @@ -311,8 +311,17 @@ def forward_dsa_prepare_npu( q_nope_out = q_nope_out.transpose(0, 1) + if enable_prefill_cp(forward_batch, m.nsa_enable_prefill_cp): + positions = cp_split_and_rebuild_position(forward_batch, positions) + q_pe, k_pe = m.rotary_emb(positions, q_pe, k_pe) + if enable_prefill_cp(forward_batch, m.nsa_enable_prefill_cp): + # support allgather+rerrange + k_nope, k_pe = m.rebuild_cp_kv_cache( + latent_cache, forward_batch, k_nope, k_pe + ) + topk_indices = m.indexer( hidden_states, q_lora, positions, forward_batch, m.layer_id ) diff --git a/python/sglang/srt/hardware_backend/npu/utils.py b/python/sglang/srt/hardware_backend/npu/utils.py index 97c126db5..070f1b252 100644 --- a/python/sglang/srt/hardware_backend/npu/utils.py +++ b/python/sglang/srt/hardware_backend/npu/utils.py @@ -68,6 +68,14 @@ def init_npu_backend(): assert _is_npu, "NPU backend initialization called on non-NPU device." import sgl_kernel_npu # noqa: F401 + + try: + import custom_ops # noqa: F401 + except ImportError: + logger.warning( + f"custom_ops not found, dsv3.2 requires this package, which includes the npu_lightning_indexer and npu_sparse_flash_attention operators." + ) + import torch_npu from torch_npu.contrib import transfer_to_npu # noqa: F401 diff --git a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py index 307f1979a..c89c38f5d 100644 --- a/python/sglang/srt/layers/attention/nsa/nsa_indexer.py +++ b/python/sglang/srt/layers/attention/nsa/nsa_indexer.py @@ -23,11 +23,7 @@ from sglang.srt.layers.attention.nsa.utils import ( cp_all_gather_rerange_output, is_nsa_enable_prefill_cp, ) -from sglang.srt.layers.dp_attention import ( - get_attention_tp_group, - get_attention_tp_rank, - get_attention_tp_size, -) +from sglang.srt.layers.dp_attention import get_attention_tp_rank, get_attention_tp_size from sglang.srt.layers.linear import ReplicatedLinear from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.rotary_embedding import get_rope_wrapper @@ -965,24 +961,12 @@ class Indexer(CustomOp): forward_batch: ForwardBatch, layer_id: int, ) -> torch.Tensor: - import custom_ops # noqa: F401 - import torch_npu - - from sglang.srt.layers.dp_attention import ( - get_attention_tp_rank, - get_attention_tp_size, - ) - from sglang.srt.utils import get_bool_env_var - if forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int is None: actual_seq_lengths_kv = forward_batch.attn_backend.forward_metadata.seq_lens else: actual_seq_lengths_kv = ( forward_batch.attn_backend.forward_metadata.seq_lens_cpu_int ) - enable_index_cp = ( - get_bool_env_var("SGLANG_USE_AG_AFTER_QLORA") and layer_id >= 4 - ) is_prefill = ( forward_batch.forward_mode.is_extend() and not forward_batch.forward_mode.is_draft_extend_v2() @@ -990,31 +974,12 @@ class Indexer(CustomOp): and not forward_batch.forward_mode.is_draft_extend() ) - attention_tp_rank = get_attention_tp_rank() - attention_tp_size = get_attention_tp_size() - cos_sin = self.rotary_emb.cos_sin_cache[positions] cos, sin = cos_sin.chunk(2, dim=-1) cos = cos.repeat(1, 2).view(-1, 1, 1, self.rope_head_dim) sin = sin.repeat(1, 2).view(-1, 1, 1, self.rope_head_dim) - if is_prefill and enable_index_cp: - slice_length = cos.shape[0] // attention_tp_size - cos = cos[ - slice_length - * attention_tp_rank : slice_length - * (attention_tp_rank + 1) - ] - sin = sin[ - slice_length - * attention_tp_rank : slice_length - * (attention_tp_rank + 1) - ] - - slot_mapping = forward_batch.out_cache_loc - block_table = forward_batch.attn_backend.forward_metadata.block_tables bs = x.shape[0] - q = self.wq_b(q_lora)[0] # [bs, 1536] @ [1536, 64 * 128] = [bs, 64 * 128] q = q.view(bs, self.n_heads, self.head_dim) # [bs, 64, 128] q_pe, q_nope = torch.split( @@ -1024,7 +989,7 @@ class Indexer(CustomOp): ) # [bs, 64, 64 + 64] q_pe = q_pe.view(bs, self.n_heads, 1, self.rope_head_dim) - q_pe = torch_npu.npu_rotary_mul(q_pe, cos, sin).view( + q_pe = torch.ops.npu.npu_rotary_mul(q_pe, cos, sin).view( bs, self.n_heads, self.rope_head_dim ) # [bs, n, d] q = torch.cat([q_pe, q_nope], dim=-1) @@ -1038,45 +1003,45 @@ class Indexer(CustomOp): ) # [bs, 64 + 64] k_pe = k_pe.view(-1, 1, 1, self.rope_head_dim) - k_pe = torch_npu.npu_rotary_mul(k_pe, cos, sin).view( + k_pe = torch.ops.npu.npu_rotary_mul(k_pe, cos, sin).view( bs, 1, self.rope_head_dim ) # [bs, 1, d] k = torch.cat([k_pe, k_nope.unsqueeze(1)], dim=-1) # [bs, 1, 128] - if is_prefill and enable_index_cp: - k, local_k = ( - torch.empty( - (k.shape[0] * attention_tp_size, k.shape[1], k.shape[2]), - dtype=k.dtype, - device=k.device, - ), - k, + if ( + is_prefill + and self.nsa_enable_prefill_cp + and forward_batch.nsa_cp_metadata is not None + ): + k = cp_all_gather_rerange_output( + k.contiguous().view(-1, self.head_dim), + self.cp_size, + forward_batch, + torch.npu.current_stream(), ) - get_attention_tp_group().all_gather_into_tensor(k, local_k) - forward_batch.token_to_kv_pool.set_index_k_buffer(layer_id, slot_mapping, k) - - indexer_input = {} + forward_batch.token_to_kv_pool.set_index_k_buffer( + layer_id, forward_batch.out_cache_loc, k + ) if is_prefill: - actual_seq_lengths_kv = forward_batch.seq_lens.to(device=q.device) - actual_seq_lengths_q = forward_batch.seq_lens.cumsum(dim=0).to( - device=q.device - ) - if enable_index_cp: - actual_seq_lengths_q -= bs * attention_tp_rank - actual_seq_lengths_q = torch.max( - actual_seq_lengths_q, - torch.zeros_like(actual_seq_lengths_q).to( - device=actual_seq_lengths_q.device - ), + if self.nsa_enable_prefill_cp and forward_batch.nsa_cp_metadata is not None: + forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q = ( + forward_batch.nsa_cp_metadata.actual_seq_q_prev_tensor, + forward_batch.nsa_cp_metadata.actual_seq_q_next_tensor, ) - actual_seq_lengths_q = torch.min( - actual_seq_lengths_q, - torch.full(actual_seq_lengths_q.shape, bs).to( - device=actual_seq_lengths_q.device - ), + forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv = ( + forward_batch.nsa_cp_metadata.kv_len_prev_tensor, + forward_batch.nsa_cp_metadata.kv_len_next_tensor, ) - + actual_seq_lengths_q = ( + forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q + ) + actual_seq_lengths_kv = ( + forward_batch.attn_backend.forward_metadata.actual_seq_lengths_kv + ) + else: + actual_seq_lengths_kv = forward_batch.seq_lens + actual_seq_lengths_q = forward_batch.seq_lens.cumsum(dim=0) else: if forward_batch.attn_backend.forward_metadata.actual_seq_lengths_q is None: if ( @@ -1109,38 +1074,96 @@ class Indexer(CustomOp): x = x.view(-1, self.hidden_size) weights = self.weights_proj(x.float())[0].to(torch.bfloat16) - block_table = ( - block_table[: actual_seq_lengths_q.size()[0]] if is_prefill else block_table - ) + block_table = forward_batch.attn_backend.forward_metadata.block_tables + if ( + is_prefill + and self.nsa_enable_prefill_cp + and forward_batch.nsa_cp_metadata is not None + ): + block_table = block_table[: actual_seq_lengths_q[0].numel()] + topk_indices = self.do_npu_cp_balance_indexer( + q.view(-1, self.n_heads, self.head_dim), + past_key_states, + weights, + actual_seq_lengths_q, + actual_seq_lengths_kv, + block_table, + ) + else: + block_table = ( + block_table[: actual_seq_lengths_q.size()[0]] + if is_prefill + else block_table + ) - topk_indices = torch.ops.custom.npu_lightning_indexer( - query=q.view(-1, self.n_heads, self.head_dim), + topk_indices = torch.ops.custom.npu_lightning_indexer( + query=q.view(-1, self.n_heads, self.head_dim), + key=past_key_states, + weights=weights, + actual_seq_lengths_query=actual_seq_lengths_q.to(torch.int32), + actual_seq_lengths_key=actual_seq_lengths_kv.to(k.device).to( + torch.int32 + ), + block_table=block_table, + layout_query="TND", + layout_key="PA_BSND", + sparse_count=self.index_topk, + sparse_mode=3, + ) + + return topk_indices + + def do_npu_cp_balance_indexer( + self, + q, + past_key_states, + indexer_weights, + actual_seq_lengths_q, + actual_seq_lengths_kv, + block_table, + ): + q_prev, q_next = torch.split(q, (q.size(0) + 1) // 2, dim=0) + weights_prev, weights_next = None, None + if indexer_weights is not None: + weights_prev, weights_next = torch.split( + indexer_weights, (indexer_weights.size(0) + 1) // 2, dim=0 + ) + weights_prev = weights_prev.contiguous().view(-1, weights_prev.shape[-1]) + weights_next = weights_next.contiguous().view(-1, weights_next.shape[-1]) + + actual_seq_lengths_q_prev, actual_seq_lengths_q_next = actual_seq_lengths_q + actual_seq_lengths_kv_prev, actual_seq_lengths_kv_next = actual_seq_lengths_kv + + topk_indices_prev = torch.ops.custom.npu_lightning_indexer( + query=q_prev, key=past_key_states, - weights=weights, - actual_seq_lengths_query=actual_seq_lengths_q.to(torch.int32), - actual_seq_lengths_key=actual_seq_lengths_kv.to(k.device).to(torch.int32), + weights=weights_prev, + actual_seq_lengths_query=actual_seq_lengths_q_prev.to( + device=q.device, dtype=torch.int32 + ), + actual_seq_lengths_key=actual_seq_lengths_kv_prev.to( + device=q.device, dtype=torch.int32 + ), block_table=block_table, layout_query="TND", layout_key="PA_BSND", sparse_count=self.index_topk, sparse_mode=3, ) - - if is_prefill and enable_index_cp: - topk_indices, local_topk_indices = ( - torch.empty( - ( - topk_indices.shape[0] * attention_tp_size, - topk_indices.shape[1], - topk_indices.shape[2], - ), - dtype=topk_indices.dtype, - device=topk_indices.device, - ), - topk_indices, - ) - get_attention_tp_group().all_gather_into_tensor( - topk_indices, local_topk_indices - ) - - return topk_indices + topk_indices_next = torch.ops.custom.npu_lightning_indexer( + query=q_next, + key=past_key_states, + weights=weights_next, + actual_seq_lengths_query=actual_seq_lengths_q_next.to( + device=q.device, dtype=torch.int32 + ), + actual_seq_lengths_key=actual_seq_lengths_kv_next.to( + device=q.device, dtype=torch.int32 + ), + block_table=block_table, + layout_query="TND", + layout_key="PA_BSND", + sparse_count=self.index_topk, + sparse_mode=3, + ) + return topk_indices_prev, topk_indices_next diff --git a/python/sglang/srt/layers/attention/nsa/utils.py b/python/sglang/srt/layers/attention/nsa/utils.py index 37817ce88..24371cd4e 100644 --- a/python/sglang/srt/layers/attention/nsa/utils.py +++ b/python/sglang/srt/layers/attention/nsa/utils.py @@ -49,6 +49,10 @@ class NSAContextParallelMetadata: kv_len_next: int = -1 actual_seq_q_prev: int = -1 actual_seq_q_next: int = -1 + kv_len_prev_tensor: torch.Tensor = None + kv_len_next_tensor: torch.Tensor = None + actual_seq_q_prev_tensor: torch.Tensor = None + actual_seq_q_next_tensor: torch.Tensor = None total_seq_lens: torch.Tensor = None @@ -312,6 +316,19 @@ def prepare_input_dp_with_cp_dsa( # TODO Support multi-batch-cp-split, multi-batch-cp support has accuracy issues # cp_seq_index = calculate_cp_seq_idx(split_list[:], seqs_len[:]) + kv_len_prev = prefix_sum_list[cp_rank] + kv_len_next = prefix_sum_list[cp_size * 2 - cp_rank - 1] + actual_seq_q_prev = split_list[cp_rank] + actual_seq_q_next = split_list[cp_size * 2 - cp_rank - 1] + kv_len_prev_tensor = torch.tensor(kv_len_prev).to(device="cuda", dtype=torch.int32) + kv_len_next_tensor = torch.tensor(kv_len_next).to(device="cuda", dtype=torch.int32) + actual_seq_q_prev_tensor = torch.tensor(actual_seq_q_prev).to( + device="cuda", dtype=torch.int32 + ) + actual_seq_q_next_tensor = torch.tensor(actual_seq_q_next).to( + device="cuda", dtype=torch.int32 + ) + nsa_cp_metadata = NSAContextParallelMetadata( split_list=split_list, max_rank_len=max_rank_len, @@ -319,10 +336,14 @@ def prepare_input_dp_with_cp_dsa( per_rank_actual_token=per_rank_actual_token, reverse_split_len=reverse_split_len, cp_reverse_index=cp_reverse_index, - kv_len_prev=prefix_sum_list[cp_rank], - kv_len_next=prefix_sum_list[cp_size * 2 - cp_rank - 1], - actual_seq_q_prev=split_list[cp_rank], - actual_seq_q_next=split_list[cp_size * 2 - cp_rank - 1], + kv_len_prev=kv_len_prev, + kv_len_next=kv_len_next, + actual_seq_q_prev=actual_seq_q_prev, + actual_seq_q_next=actual_seq_q_next, + kv_len_prev_tensor=kv_len_prev_tensor, + kv_len_next_tensor=kv_len_next_tensor, + actual_seq_q_prev_tensor=actual_seq_q_prev_tensor, + actual_seq_q_next_tensor=actual_seq_q_next_tensor, total_seq_lens=kv_len_origin, ) return nsa_cp_metadata diff --git a/python/sglang/srt/layers/communicator_nsa_cp.py b/python/sglang/srt/layers/communicator_nsa_cp.py index 643433865..2b323edf1 100644 --- a/python/sglang/srt/layers/communicator_nsa_cp.py +++ b/python/sglang/srt/layers/communicator_nsa_cp.py @@ -176,14 +176,13 @@ class NSACPCommunicateWithAllReduceAndLayerNormFn( *, residual_input_mode, ): - if context.attn_dp_size != 1: - if nsa_enable_prefill_cp(): - hidden_states += residual - if hidden_states.shape[0] != 0: - hidden_states = layernorm(hidden_states) - return hidden_states, residual - else: - assert False, "not yet handled" + if nsa_enable_prefill_cp(): + hidden_states += residual + if hidden_states.shape[0] != 0: + hidden_states = layernorm(hidden_states) + return hidden_states, residual + else: + assert False, "not yet handled" @staticmethod def _scatter_hidden_states_and_residual( diff --git a/test/srt/ascend/test_ascend_tp4_bf16.py b/test/srt/ascend/test_ascend_tp4_bf16.py index e1feac417..9730ffeb3 100644 --- a/test/srt/ascend/test_ascend_tp4_bf16.py +++ b/test/srt/ascend/test_ascend_tp4_bf16.py @@ -36,6 +36,7 @@ class TestAscendTp4Bf16(CustomTestCase): 32, "--attention-backend", "ascend", + "--disable-radix-cache", "--cuda-graph-max-bs", 32, "--tp-size",