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(
|
||||
|
||||
Reference in New Issue
Block a user