Speculative Decoding support for intel_xpu attention backend on XPU target (#30548)
This commit is contained in:
@@ -0,0 +1,57 @@
|
||||
"""intel_xpu attention backend (EAGLE3 topk=1 chain + EAGLE/Llama-2 spec)."""
|
||||
|
||||
import unittest
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.test.ci.ci_register import register_xpu_ci
|
||||
from sglang.test.kits.matched_stop_kit import MatchedStopMixin
|
||||
from sglang.test.kits.spec_server_kits import (
|
||||
SpecAccuracyKit,
|
||||
SpecFeatureKit,
|
||||
SpecHiddenStatesKit,
|
||||
SpecLogprobKit,
|
||||
SpecPenaltyKit,
|
||||
)
|
||||
from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base
|
||||
|
||||
register_xpu_ci(est_time=1800, suite="nightly-xpu-1-gpu", nightly=True)
|
||||
|
||||
|
||||
class TestEagle3IntelXPU(
|
||||
Eagle3Base,
|
||||
MatchedStopMixin,
|
||||
SpecAccuracyKit,
|
||||
SpecLogprobKit,
|
||||
SpecPenaltyKit,
|
||||
SpecFeatureKit,
|
||||
):
|
||||
"""EAGLE3 spec v2 on the intel_xpu attention backend (kits listed in bases)."""
|
||||
|
||||
attention_backend = "intel_xpu"
|
||||
max_running_requests = 24
|
||||
gsm8k_num_examples = 300
|
||||
gsm8k_check_accept_len = True
|
||||
mem_fraction_static = 0.95
|
||||
extra_args = ("--max-total-tokens", "16384", "--disable-decode-cuda-graph")
|
||||
env_overrides = ((envs.SGLANG_ENABLE_STRICT_MEM_CHECK_DURING_BUSY, 1),)
|
||||
|
||||
|
||||
class TestEagleLlama2IntelXPU(
|
||||
EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit, SpecHiddenStatesKit
|
||||
):
|
||||
"""EAGLE/Llama-2 on intel_xpu using the supported topk = 1 paged config."""
|
||||
|
||||
attention_backend = "intel_xpu"
|
||||
spec_topk = 1
|
||||
page_size = 64
|
||||
gsm8k_check_accept_len = True
|
||||
gsm8k_num_examples = 300
|
||||
enable_return_hidden_states = True
|
||||
mem_fraction_static = 0.95
|
||||
max_running_requests = 6
|
||||
chunked_prefill_size = 512
|
||||
extra_args = ("--max-total-tokens", "16384", "--disable-decode-cuda-graph")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -57,8 +57,8 @@ class TestEncoderDecoderForward(unittest.TestCase):
|
||||
# The caller picks (page_table, cache_seqlens, causal) via
|
||||
# _encoder_decoder_page_table -- cross-attn -> encoder_page_table +
|
||||
# encoder_lens_int32 + causal=False; self-attn -> page_table +
|
||||
# cache_seqlens_int32 + causal=True -- then hands them to the generic
|
||||
# _forward_attn_flat_page_table, which must forward them unchanged with a
|
||||
# cache_seqlens_int32 + causal=True -- then hands them to
|
||||
# _forward_encoder_decoder_attn, which must forward them unchanged with a
|
||||
# page_size=1 k_cache (shape[1]==1) so PR #454 routes to the varlen gather.
|
||||
enc_pt = torch.arange(5, dtype=torch.int32).unsqueeze(0)
|
||||
dec_pt = (torch.arange(4, dtype=torch.int32) + 10).unsqueeze(0)
|
||||
@@ -70,7 +70,7 @@ class TestEncoderDecoderForward(unittest.TestCase):
|
||||
)
|
||||
key_cache = self.k_flat.view(-1, 1, self.HK, self.D)
|
||||
value_cache = self.v_flat.view(-1, 1, self.HK, self.D)
|
||||
q = torch.randn(1, self.HQ * self.D)
|
||||
q = torch.randn(1, self.HQ, self.D)
|
||||
cu_seqlens_q = torch.tensor([0, 1], dtype=torch.int32)
|
||||
|
||||
for is_cross, exp_pt, exp_seqlens, exp_causal in (
|
||||
@@ -94,15 +94,16 @@ class TestEncoderDecoderForward(unittest.TestCase):
|
||||
)
|
||||
|
||||
with patch.object(xpu_backend, "flash_attn_with_kvcache", fake_kvcache):
|
||||
self.backend._forward_attn_flat_page_table(
|
||||
self.backend._forward_encoder_decoder_attn(
|
||||
q=q,
|
||||
key_cache=key_cache,
|
||||
value_cache=value_cache,
|
||||
layer=layer,
|
||||
page_table=page_table,
|
||||
cache_seqlens=cache_seqlens,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
max_seqlen_q=1,
|
||||
scale=layer.scaling,
|
||||
softcap=layer.logit_cap,
|
||||
causal=causal,
|
||||
)
|
||||
self.assertTrue(torch.equal(captured["page_table"], exp_pt))
|
||||
@@ -117,18 +118,20 @@ class TestEncoderDecoderForward(unittest.TestCase):
|
||||
# zeros and never launch the kernel.
|
||||
key_cache = self.k_flat.view(-1, 1, self.HK, self.D)
|
||||
value_cache = self.v_flat.view(-1, 1, self.HK, self.D)
|
||||
q = torch.randn(1, self.HQ * self.D)
|
||||
q = torch.randn(1, self.HQ, self.D)
|
||||
sentinel = MagicMock(side_effect=AssertionError("kernel must not run"))
|
||||
layer = self._layer(True)
|
||||
with patch.object(xpu_backend, "flash_attn_with_kvcache", sentinel):
|
||||
out = self.backend._forward_attn_flat_page_table(
|
||||
out = self.backend._forward_encoder_decoder_attn(
|
||||
q=q,
|
||||
key_cache=key_cache,
|
||||
value_cache=value_cache,
|
||||
layer=self._layer(True),
|
||||
page_table=torch.zeros(1, 0, dtype=torch.int32),
|
||||
cache_seqlens=torch.zeros(1, dtype=torch.int32),
|
||||
cu_seqlens_q=torch.tensor([0, 1], dtype=torch.int32),
|
||||
max_seqlen_q=1,
|
||||
scale=layer.scaling,
|
||||
softcap=layer.logit_cap,
|
||||
causal=False,
|
||||
)
|
||||
sentinel.assert_not_called()
|
||||
@@ -141,7 +144,7 @@ class TestEncoderDecoderForward(unittest.TestCase):
|
||||
# counts (2 and 3) exercise the cu_seqlens_q -> per-request row mapping.
|
||||
key_cache = self.k_flat.view(-1, 1, self.HK, self.D)
|
||||
value_cache = self.v_flat.view(-1, 1, self.HK, self.D)
|
||||
q = torch.randn(5, self.HQ * self.D)
|
||||
q = torch.randn(5, self.HQ, self.D)
|
||||
|
||||
def fake_kvcache(*_, **kw):
|
||||
# All-ones (never-NaN) sentinel so zeroed rows are distinguishable.
|
||||
@@ -149,16 +152,18 @@ class TestEncoderDecoderForward(unittest.TestCase):
|
||||
(kw["q"].shape[0], kw["q"].shape[1], kw["v_cache"].shape[-1])
|
||||
)
|
||||
|
||||
layer = self._layer(True)
|
||||
with patch.object(xpu_backend, "flash_attn_with_kvcache", fake_kvcache):
|
||||
out = self.backend._forward_attn_flat_page_table(
|
||||
out = self.backend._forward_encoder_decoder_attn(
|
||||
q=q,
|
||||
key_cache=key_cache,
|
||||
value_cache=value_cache,
|
||||
layer=self._layer(True),
|
||||
page_table=torch.zeros(2, 4, dtype=torch.int32),
|
||||
cache_seqlens=torch.tensor([0, 4], dtype=torch.int32),
|
||||
cu_seqlens_q=torch.tensor([0, 2, 5], dtype=torch.int32),
|
||||
max_seqlen_q=3,
|
||||
scale=layer.scaling,
|
||||
softcap=layer.logit_cap,
|
||||
causal=False,
|
||||
)
|
||||
self.assertTrue(torch.equal(out[:2], torch.zeros(2, self.HQ, self.D)))
|
||||
|
||||
@@ -2,13 +2,12 @@
|
||||
|
||||
The backend calls flash_attn_with_kvcache with a page_size=1 view; sgl-kernel-xpu
|
||||
PR #454 detects that and gathers + runs varlen inside the kernel. This runs on an
|
||||
actual XPU and guards what a mocked CPU test cannot: _forward_attn_flat_page_table
|
||||
actual XPU and guards what a mocked CPU test cannot: _forward_encoder_decoder_attn
|
||||
plus the real kernel produce correct attention for a scattered (non-page-aligned)
|
||||
token-slot layout, for both cross-attn (non-causal) and decoder self-attn (causal).
|
||||
"""
|
||||
|
||||
import unittest
|
||||
from types import SimpleNamespace
|
||||
|
||||
import torch
|
||||
|
||||
@@ -85,30 +84,23 @@ class TestXPUEncoderDecoderVarlen(CustomTestCase):
|
||||
.to(self.dev)
|
||||
)
|
||||
q = torch.randn(num_rows, self.H, self.D, dtype=torch.bfloat16, device=self.dev)
|
||||
layer = SimpleNamespace(
|
||||
is_cross_attention=not causal,
|
||||
tp_q_head_num=self.H,
|
||||
tp_k_head_num=self.H,
|
||||
tp_v_head_num=self.H,
|
||||
head_dim=self.D,
|
||||
scaling=0.5,
|
||||
logit_cap=0.0,
|
||||
)
|
||||
scale, softcap = 0.5, 0.0
|
||||
key_cache = self.k_flat.view(-1, 1, self.H, self.D)
|
||||
value_cache = self.v_flat.view(-1, 1, self.H, self.D)
|
||||
|
||||
# causal=True mirrors decoder self-attn, causal=False cross-attn; the
|
||||
# generic helper takes the (page_table, cache_seqlens, causal) that the
|
||||
# caller's _encoder_decoder_page_table dispatch would have selected.
|
||||
got = self.backend._forward_attn_flat_page_table(
|
||||
got = self.backend._forward_encoder_decoder_attn(
|
||||
q=q,
|
||||
key_cache=key_cache,
|
||||
value_cache=value_cache,
|
||||
layer=layer,
|
||||
page_table=page_table,
|
||||
cache_seqlens=cache_seqlens,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
max_seqlen_q=1,
|
||||
scale=scale,
|
||||
softcap=softcap,
|
||||
causal=causal,
|
||||
)
|
||||
torch.xpu.synchronize()
|
||||
@@ -119,7 +111,7 @@ class TestXPUEncoderDecoderVarlen(CustomTestCase):
|
||||
page_table=page_table,
|
||||
cache_seqlens=cache_seqlens,
|
||||
cu_seqlens_q=cu_seqlens_q,
|
||||
scale=layer.scaling,
|
||||
scale=scale,
|
||||
causal=causal,
|
||||
)
|
||||
self.assertEqual(tuple(got.shape), (num_rows, self.H, self.D))
|
||||
|
||||
Reference in New Issue
Block a user