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:
co-authored by
github-actions[bot]
Singh
Singh
Pramod Kumar
parent
503e36cbae
commit
e24e31efd9
@@ -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()
|
||||
Reference in New Issue
Block a user