diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 37a8061b1..301c45a48 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -1527,9 +1527,19 @@ class TritonAttnBackend(AttentionBackend): dtype=torch.float32, ) - # Current chunk K/V is still local before masked cache write, so it can - # use the original extend kernel's current-token stage directly. + # Select the replicated K/V heads matching this rank's Q shard. if k.numel() > 0: + if layer.tp_k_head_num > 1: + kv_head_start = ( + group.rank_in_group * layer.tp_k_head_num // group.world_size + ) + kv_head_end = max( + (group.rank_in_group + 1) * layer.tp_k_head_num // group.world_size, + kv_head_start + 1, + ) + k = k[:, kv_head_start:kv_head_end] + v = v[:, kv_head_start:kv_head_end] + empty_kv_indptr = torch.zeros_like(kv_indptr) self.extend_attention_fwd( q_local, diff --git a/test/registered/dcp/test_dcp_layout_unit.py b/test/registered/dcp/test_dcp_layout_unit.py index 622945904..cefc57621 100644 --- a/test/registered/dcp/test_dcp_layout_unit.py +++ b/test/registered/dcp/test_dcp_layout_unit.py @@ -21,6 +21,7 @@ import torch from sglang.srt import runtime_context as rc from sglang.srt.configs.model_config import ModelConfig +from sglang.srt.layers.attention.triton_backend import TritonAttnBackend from sglang.srt.layers.dcp.layout import get_dcp_lens from sglang.srt.layers.linear import QKVParallelLinear from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator @@ -96,6 +97,77 @@ class TestGetDcpLens(CustomTestCase): lens = torch.tensor(LENS, dtype=torch.int32) self.assertTrue(torch.equal(get_dcp_lens(lens, 1, 0), lens)) + def test_gqa_current_chunk_selects_kv_for_the_global_dcp_head_layout(self): + """A local Q shard must not restart GQA mapping at KV head zero.""" + + class FakeDcpGroup: + world_size = 4 + rank_in_group = 1 + + def __init__(self): + self.all_gather_calls = 0 + + def all_gather(self, tensor, dim): + self.all_gather_calls += 1 + return torch.cat((tensor, tensor + 10), dim=dim) + + group = FakeDcpGroup() + backend = TritonAttnBackend.__new__(TritonAttnBackend) + backend.forward_metadata = SimpleNamespace( + custom_mask=None, + kv_indptr=torch.zeros(2, dtype=torch.int32), + kv_indices=torch.empty(0, dtype=torch.int64), + max_extend_len=1, + qo_indptr=torch.tensor([0, 1], dtype=torch.int64), + ) + backend.token_to_kv_pool = SimpleNamespace( + get_key_buffer=lambda _layer_id: torch.empty(0), + get_value_buffer=lambda _layer_id: torch.empty(0), + ) + + kernel_q_shapes = [] + kernel_k = [] + + def fake_extend_attention(q, k, _v, out, *_args, lse_extend, **_kwargs): + kernel_q_shapes.append(q.shape) + kernel_k.append(k.clone()) + out.copy_(q.float()) + lse_extend.zero_() + + backend.extend_attention_fwd = fake_extend_attention + layer = SimpleNamespace( + sliding_window_size=-1, + tp_q_head_num=2, + tp_k_head_num=2, + qk_head_dim=2, + v_head_dim=2, + k_scale=None, + v_scale=None, + layer_id=0, + scaling=1.0, + xai_temperature_len=-1, + ) + q = torch.arange(4, dtype=torch.bfloat16).view(1, 4) + k = torch.tensor([[[0.0, 1.0], [10.0, 11.0]]]) + + with rc.get_parallel().override(dcp_group=group): + out = backend._forward_extend_dcp( + q=q, + k=k, + v=k.clone(), + layer=layer, + forward_batch=SimpleNamespace(), + causal=True, + logits_soft_cap=0.0, + sinks=None, + ) + + self.assertEqual(group.all_gather_calls, 0) + self.assertEqual(kernel_q_shapes, [torch.Size([1, 2, 2])]) + # In TP4/DCP4 with two KV heads, ranks 0 and 1 both belong to KV head 0. + self.assertTrue(torch.equal(kernel_k[0], k[:, 0:1])) + self.assertTrue(torch.equal(out, q)) + def test_paged_allocator_exposes_dcp_virtual_capacity(self): real_kv_size = 1024 dcp_size = 4