Fix/whisper xpu varlen encoder decoder (#36298)

Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
Co-authored-by: Singh <rohitsi2@iil-login.iind.intel.com>
Co-authored-by: Singh <rohitsi2@iil-gnrap02.iind.intel.com>
Co-authored-by: Pramod Kumar <144990617+pramodkumar-habanalabs@users.noreply.github.com>
This commit is contained in:
Rohit Kumar Singh
2026-09-07 13:21:29 +08:00
committed by GitHub
co-authored by github-actions[bot] Singh Singh Pramod Kumar
parent 503e36cbae
commit e24e31efd9
3 changed files with 500 additions and 11 deletions
+121 -11
View File
@@ -110,6 +110,21 @@ class XPUAttentionBackend(AttentionBackend):
1 if get_exec().deterministic.enable_deterministic_inference else 0
)
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
if self.is_encoder_decoder:
from sglang.srt.model_executor.cuda_graph_config import (
cuda_graph_fully_disabled,
)
# Encoder-decoder cross-/self-attention below uses a dynamic-shape
# varlen KV gather (page_size=1 semantics) that cannot be captured.
# XPU disables CUDA graph by default, so this holds; the guard fails
# loudly if a future XPU graph path is force-enabled instead of
# silently mis-indexing through the paged graph-metadata path.
assert cuda_graph_fully_disabled(), (
"Encoder-decoder models (e.g. Whisper) on the intel_xpu attention "
"backend require CUDA graph disabled (off by default on XPU); the "
"graph decode path cannot run the varlen KV gather."
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Initialize forward metadata hence all layers in the forward pass can reuse it."""
@@ -361,10 +376,6 @@ class XPUAttentionBackend(AttentionBackend):
# Encoder metadata for cross attention
if forward_batch.encoder_lens is not None:
assert forward_batch.encoder_lens.numel() == 1, (
"Only encoder size 1 is supported for now"
)
metadata.encoder_lens_int32 = forward_batch.encoder_lens.to(torch.int32)
metadata.encoder_cu_seqlens_k = torch.nn.functional.pad(
torch.cumsum(metadata.encoder_lens_int32, dim=0, dtype=torch.int32),
@@ -375,12 +386,18 @@ class XPUAttentionBackend(AttentionBackend):
forward_batch.req_pool_indices, : metadata.encoder_max_seq_len_k
]
# Currently only support forward_batch.encoder_lens.numel() == 1
# Decoder self-attn KV: per-request token-granular slice starting at
# each request's own encoder offset encoder_lens[i], not a single max.
text_max = metadata.max_seq_len_k
arange_text = torch.arange(
text_max, device=forward_batch.req_pool_indices.device
)
text_col = forward_batch.encoder_lens.long().unsqueeze(
1
) + arange_text.unsqueeze(0)
text_row = forward_batch.req_pool_indices.unsqueeze(1).expand(-1, text_max)
metadata.page_table = self.req_to_token_pool.req_to_token[
forward_batch.req_pool_indices,
metadata.encoder_max_seq_len_k : (
metadata.encoder_max_seq_len_k + metadata.max_seq_len_k
),
text_row, text_col
]
# Translate full-pool indices to SWA-pool indices for hybrid models
@@ -415,8 +432,10 @@ class XPUAttentionBackend(AttentionBackend):
workspace_size, device=self.device, dtype=torch.uint8
)
# Convert the page table to a strided format which is needed by FA3 API
if self.page_size > 1:
# Convert the page table to a strided format which is needed by FA3 API.
# Encoder-decoder page_table holds token-slot indices for the varlen
# kernel (page_size=1 semantics), so it must not be page-strided.
if self.page_size > 1 and forward_batch.encoder_lens is None:
self.strided_indices = torch.arange(
0, metadata.page_table.shape[1], self.page_size, device=self.device
)
@@ -572,6 +591,22 @@ class XPUAttentionBackend(AttentionBackend):
value_cache = value_cache.view(
-1, self.page_size, layer.tp_v_head_num, layer.head_dim
)
if self.is_encoder_decoder and forward_batch.encoder_lens is not None:
page_table, cache_seqlens, causal = self._encoder_decoder_page_table(
layer, metadata
)
o = self._forward_attn_flat_page_table(
q=q,
key_cache=key_cache,
value_cache=value_cache,
layer=layer,
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=metadata.cu_seqlens_q,
max_seqlen_q=metadata.max_seq_len_q,
causal=causal,
)
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
if layer.is_cross_attention:
page_table = metadata.encoder_page_table
cache_seqlens = metadata.encoder_lens_int32
@@ -765,6 +800,64 @@ class XPUAttentionBackend(AttentionBackend):
out = o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
return out
@staticmethod
def _encoder_decoder_page_table(layer, metadata):
"""Pick (page_table, cache_seqlens, causal) for an encoder-decoder layer:
cross-attention reads the encoder KV region (non-causal), decoder
self-attention reads the decoder KV region (causal)."""
if layer.is_cross_attention:
return metadata.encoder_page_table, metadata.encoder_lens_int32, False
return metadata.page_table, metadata.cache_seqlens_int32, True
def _forward_attn_flat_page_table(
self,
*,
q,
key_cache,
value_cache,
layer,
page_table,
cache_seqlens,
cu_seqlens_q,
max_seqlen_q,
causal,
):
"""MHA on XPU via flash_attn_with_kvcache with a page_size=1 (flat
token-slot) page table. sgl-kernel-xpu PR #454 detects a page_size==1
k_cache + page_table and gathers + runs varlen internally, so the backend
calls it like the FA (CUDA) path. Eager-only. A request with
cache_seqlens==0 attends to no keys and the kernel returns NaN for its
rows, so those rows are zeroed (an all-empty batch skips the launch).
"""
q_rows = q.contiguous().view(-1, layer.tp_q_head_num, layer.head_dim)
# If the maximum cache_seqlens is 0, there are no keys to attend to.
if int(cache_seqlens.max().item()) == 0:
return q_rows.new_zeros(
(q_rows.shape[0], q_rows.shape[1], value_cache.shape[-1])
)
# page_size=1 view of the KV pool: PR #454 detects if page size dim is 1
# i.e. k_cache.shape[1] == 1 and routes flash_attn_with_kvcache to varlen gather.
k_cache = key_cache.reshape(-1, 1, layer.tp_k_head_num, layer.head_dim)
v_cache = value_cache.reshape(-1, 1, layer.tp_v_head_num, layer.head_dim)
out = flash_attn_with_kvcache(
q=q_rows,
k_cache=k_cache,
v_cache=v_cache,
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=cu_seqlens_q,
max_seqlen_q=max_seqlen_q,
softmax_scale=layer.scaling,
causal=causal,
softcap=layer.logit_cap,
)
# Mixed batch: requests with cache_seqlens==0 attend to no keys and come
# back as NaN, so zero their query rows (mapped via cu_seqlens_q).
if int(cache_seqlens.min().item()) == 0:
seg = cu_seqlens_q[1:] - cu_seqlens_q[:-1]
out[(cache_seqlens == 0).repeat_interleave(seg)] = 0
return out
def forward_decode(
self,
q: torch.Tensor,
@@ -870,6 +963,23 @@ class XPUAttentionBackend(AttentionBackend):
-1, self.page_size, layer.tp_v_head_num, layer.head_dim
)
if self.is_encoder_decoder and forward_batch.encoder_lens is not None:
page_table, cache_seqlens, causal = self._encoder_decoder_page_table(
layer, metadata
)
o = self._forward_attn_flat_page_table(
q=q,
key_cache=key_cache,
value_cache=value_cache,
layer=layer,
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=metadata.cu_seqlens_q,
max_seqlen_q=1,
causal=causal,
)
return o.view(-1, layer.tp_q_head_num * layer.v_head_dim)
if layer.is_cross_attention:
# Always use non-chunked logic for cross-attention
o = flash_attn_with_kvcache(
@@ -0,0 +1,219 @@
import sys
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.test.ci.ci_register import register_cpu_ci
with patch.dict(
sys.modules,
{
module: MagicMock()
for module in (
"sgl_kernel",
"sgl_kernel.flash_attn",
"sgl_kernel.quantization",
"sgl_kernel.scalar_type",
)
},
):
from sglang.srt.layers.attention import xpu_backend
from sglang.srt.layers.attention.xpu_backend import XPUAttentionBackend
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class TestEncoderDecoderForward(unittest.TestCase):
"""Encoder-decoder attention on XPU calls flash_attn_with_kvcache with a
page_size=1 view; sgl-kernel-xpu PR #454 gathers the token-slot page_table and
runs varlen inside the kernel. These guard the backend's own responsibilities
-- the cross-vs-self dispatch, the page_size=1 view, the empty-KV zero guard,
and the per-request encoder-offset metadata. Real-kernel numerical correctness
is covered on device in test/registered/xpu/test_xpu_encoder_decoder_varlen.py.
"""
HQ, HK, D, TOTAL_SLOTS = 4, 2, 8, 40
def setUp(self):
torch.manual_seed(0)
self.backend = XPUAttentionBackend.__new__(XPUAttentionBackend)
self.k_flat = torch.randn(self.TOTAL_SLOTS, self.HK, self.D)
self.v_flat = torch.randn(self.TOTAL_SLOTS, self.HK, self.D)
def _layer(self, is_cross):
return SimpleNamespace(
is_cross_attention=is_cross,
tp_q_head_num=self.HQ,
tp_k_head_num=self.HK,
tp_v_head_num=self.HK,
head_dim=self.D,
scaling=0.5,
logit_cap=0.0,
)
def test_dispatch_and_forward_cross_vs_self(self):
# 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
# 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)
metadata = SimpleNamespace(
encoder_page_table=enc_pt,
encoder_lens_int32=torch.tensor([5], dtype=torch.int32),
page_table=dec_pt,
cache_seqlens_int32=torch.tensor([4], dtype=torch.int32),
)
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)
cu_seqlens_q = torch.tensor([0, 1], dtype=torch.int32)
for is_cross, exp_pt, exp_seqlens, exp_causal in (
(True, enc_pt, metadata.encoder_lens_int32, False),
(False, dec_pt, metadata.cache_seqlens_int32, True),
):
layer = self._layer(is_cross)
page_table, cache_seqlens, causal = (
self.backend._encoder_decoder_page_table(layer, metadata)
)
self.assertTrue(torch.equal(page_table, exp_pt))
self.assertTrue(torch.equal(cache_seqlens, exp_seqlens))
self.assertEqual(causal, exp_causal)
captured = {}
def fake_kvcache(*_, **kw):
captured.update(kw)
return kw["q"].new_zeros(
(kw["q"].shape[0], kw["q"].shape[1], kw["v_cache"].shape[-1])
)
with patch.object(xpu_backend, "flash_attn_with_kvcache", fake_kvcache):
self.backend._forward_attn_flat_page_table(
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,
causal=causal,
)
self.assertTrue(torch.equal(captured["page_table"], exp_pt))
self.assertTrue(torch.equal(captured["cache_seqlens"], exp_seqlens))
self.assertEqual(captured["causal"], exp_causal)
self.assertEqual(captured["k_cache"].shape[1], 1) # page_size=1 view
self.assertEqual(captured["max_seqlen_q"], 1)
def test_all_empty_returns_zeros_without_kernel(self):
# Whisper text-only warmup: all cache_seqlens == 0. PR #454's page_size=1
# path returns NaN for an empty KV, so the backend must short-circuit to
# 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)
sentinel = MagicMock(side_effect=AssertionError("kernel must not run"))
with patch.object(xpu_backend, "flash_attn_with_kvcache", sentinel):
out = self.backend._forward_attn_flat_page_table(
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,
causal=False,
)
sentinel.assert_not_called()
self.assertTrue(torch.equal(out, torch.zeros(1, self.HQ, self.D)))
def test_mixed_empty_zeros_only_empty_request_rows(self):
# Mixed batch: request 0 has cache_seqlens==0 (no keys), request 1 has
# keys. PR #454 returns NaN for the empty request's rows, so the backend
# must zero exactly those rows and leave the rest untouched. Unequal query
# 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)
def fake_kvcache(*_, **kw):
# All-ones (never-NaN) sentinel so zeroed rows are distinguishable.
return kw["q"].new_ones(
(kw["q"].shape[0], kw["q"].shape[1], kw["v_cache"].shape[-1])
)
with patch.object(xpu_backend, "flash_attn_with_kvcache", fake_kvcache):
out = self.backend._forward_attn_flat_page_table(
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,
causal=False,
)
self.assertTrue(torch.equal(out[:2], torch.zeros(2, self.HQ, self.D)))
self.assertTrue(torch.equal(out[2:], torch.ones(3, self.HQ, self.D)))
def test_init_forward_metadata_per_request_encoder_offset(self):
# Guards the encoder_lens.numel()==1 removal: with UNEQUAL encoder lengths,
# init_forward_metadata must place each request's decoder self-attn page
# table at ITS OWN encoder_lens[i] offset (not a single batch max), and
# slice encoder KV per request. It must also skip the //page_size stride
# for enc-dec (token-slot indices feed the varlen kernel). Fails on the old
# scalar-max-offset slice (row0 would start at col 5, not 3).
from sglang.srt.model_executor.forward_batch_info import ForwardMode
backend = XPUAttentionBackend.__new__(XPUAttentionBackend)
backend.page_size = 128 # >1: also exercises the enc-dec stride-skip
backend.is_encoder_decoder = True
backend.use_mla = False
backend.use_sliding_window_kv_pool = False
backend.attention_chunk_size = None
backend.topk = 0
# req_to_token[i, j] = 100*i + j, so gathered values reveal (row, col).
req_to_token = torch.arange(16).unsqueeze(0) + torch.tensor([[0], [100]])
backend.req_to_token_pool = SimpleNamespace(req_to_token=req_to_token)
fb = SimpleNamespace(
forward_mode=ForwardMode.DECODE,
seq_lens=torch.tensor([2, 4], dtype=torch.int64), # decoder lengths
seq_lens_cpu=torch.tensor([2, 4]),
batch_size=2,
req_pool_indices=torch.tensor([0, 1]),
encoder_lens=torch.tensor([3, 5], dtype=torch.int64), # UNEQUAL
spec_info=None,
out_cache_loc=None,
)
backend.init_forward_metadata(fb)
md = backend.forward_metadata
# Encoder KV: columns [0 : max_enc=5] of each request's row; per-request
# lengths + segment boundaries captured for the kernel's internal gather.
self.assertEqual(
md.encoder_page_table.tolist(),
[[0, 1, 2, 3, 4], [100, 101, 102, 103, 104]],
)
self.assertEqual(md.encoder_lens_int32.tolist(), [3, 5])
self.assertEqual(md.encoder_cu_seqlens_k.tolist(), [0, 3, 8])
# Decoder self-attn KV (text_max = max(seq_lens) = 4 columns each): request 0
# starts at col 3 (its encoder_len), request 1 at col 5 (its encoder_len).
# Token-granular (not //128), proving the stride transform was skipped.
self.assertEqual(
md.page_table.tolist(),
[[3, 4, 5, 6], [105, 106, 107, 108]],
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,160 @@
"""Real-device XPU tests for the encoder-decoder attention path.
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
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
from sglang.srt.layers.attention.xpu_backend import XPUAttentionBackend
from sglang.test.ci.ci_register import register_xpu_ci
from sglang.test.test_utils import CustomTestCase
register_xpu_ci(est_time=15, suite="stage-b-test-1-gpu-xpu")
def _sdpa_ref(
*, q, k_flat, v_flat, page_table, cache_seqlens, cu_seqlens_q, scale, causal
):
"""Per-request SDPA oracle: gather each request's valid slots straight from
the page_table and attend. fp32 math for a stable bf16 comparison."""
outs = []
for i in range(page_table.shape[0]):
qs, qe = int(cu_seqlens_q[i]), int(cu_seqlens_q[i + 1])
sl = int(cache_seqlens[i])
q_i = q[qs:qe].float()
slots = page_table[i, :sl].long()
k_i = k_flat.index_select(0, slots).float()
v_i = v_flat.index_select(0, slots).float()
scores = torch.einsum("qhd,khd->hqk", q_i, k_i) * scale
lq, lk = qe - qs, sl
if causal and lq > 0 and lk > 0:
row = torch.arange(lq, device=q.device).unsqueeze(1)
col = torch.arange(lk, device=q.device).unsqueeze(0)
keep = col <= (lk - lq + row)
scores = scores.masked_fill(~keep.unsqueeze(0), float("-inf"))
probs = scores.softmax(dim=-1)
outs.append(torch.einsum("hqk,khd->qhd", probs, v_i))
return torch.cat(outs, dim=0)
@unittest.skipUnless(
hasattr(torch, "xpu") and torch.xpu.is_available(), "requires an Intel XPU"
)
class TestXPUEncoderDecoderVarlen(CustomTestCase):
# Whisper-large-v3 is MHA (num_kv_heads == num_heads), so use MHA here to keep
# the reference exact; head_dim=64 satisfies the kernel's alignment.
H, D, TOTAL_SLOTS = 8, 64, 64
def setUp(self):
torch.manual_seed(0)
self.dev = torch.device("xpu")
self.backend = XPUAttentionBackend.__new__(XPUAttentionBackend)
self.backend.is_encoder_decoder = True
# Deliberately scattered (non-page-aligned) slot indices: a paged kernel
# would mis-read these; the page_size=1 gather must be alignment-agnostic.
perm = torch.randperm(self.TOTAL_SLOTS)
self.k_flat = torch.randn(
self.TOTAL_SLOTS, self.H, self.D, dtype=torch.bfloat16, device=self.dev
)[perm].contiguous()
self.v_flat = torch.randn(
self.TOTAL_SLOTS, self.H, self.D, dtype=torch.bfloat16, device=self.dev
)[perm].contiguous()
def _check(self, *, cache_seqlens, cu_seqlens_q, causal):
cache_seqlens = cache_seqlens.to(self.dev)
cu_seqlens_q = cu_seqlens_q.to(self.dev)
num_rows = int(cu_seqlens_q[-1])
m = int(cache_seqlens.max())
# Rows packed valid-first; scatter distinct slots per request (build the
# permutation on CPU, then move -- randperm(device="xpu") is unreliable).
page_table = (
torch.stack(
[
torch.randperm(self.TOTAL_SLOTS)[:m]
for _ in range(cache_seqlens.numel())
]
)
.to(torch.int32)
.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,
)
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(
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,
causal=causal,
)
torch.xpu.synchronize()
want = _sdpa_ref(
q=q,
k_flat=self.k_flat,
v_flat=self.v_flat,
page_table=page_table,
cache_seqlens=cache_seqlens,
cu_seqlens_q=cu_seqlens_q,
scale=layer.scaling,
causal=causal,
)
self.assertEqual(tuple(got.shape), (num_rows, self.H, self.D))
self.assertTrue(torch.isfinite(got).all(), "attention output must be finite")
# bf16 kernel vs fp32 reference: loose tolerance.
torch.testing.assert_close(got.float(), want, rtol=2e-2, atol=2e-2)
def test_cross_attention_decode_on_xpu(self):
# 1 query/request, attend all encoder KV, non-causal, unequal lengths.
self._check(
cache_seqlens=torch.tensor([5, 8], dtype=torch.int32),
cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32),
causal=False,
)
def test_decoder_self_attention_decode_on_xpu(self):
self._check(
cache_seqlens=torch.tensor([4, 6], dtype=torch.int32),
cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32),
causal=True,
)
def test_mixed_empty_batch_on_xpu(self):
# Mixed batch: request 0 has no keys (cache_seqlens==0), request 1 has some.
# On the real kernel the empty request's rows come back NaN/inf, so the
# backend must zero them without corrupting request 1. The SDPA oracle
# yields zeros for the empty request (empty-key contraction), so the shared
# assert_close plus the finiteness check guard against a regression that
# drops the zeroing and leaks NaN into the output.
self._check(
cache_seqlens=torch.tensor([0, 6], dtype=torch.int32),
cu_seqlens_q=torch.tensor([0, 1, 2], dtype=torch.int32),
causal=False,
)
if __name__ == "__main__":
unittest.main()