Add sliding-window mask support to TorchNativeAttnBackend (#26516)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-05-28 01:06:57 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent a245cae3d1
commit 00cd6fb3d9
@@ -24,6 +24,21 @@ class TorchNativeAttnBackend(AttentionBackend):
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool self.token_to_kv_pool = model_runner.token_to_kv_pool
@staticmethod
def _make_sliding_window_mask(
*,
q_len: int,
kv_len: int,
sliding_window_size: int,
device: torch.device,
query_offset: int = 0,
) -> torch.Tensor:
q_pos = torch.arange(
query_offset, query_offset + q_len, device=device
).unsqueeze(1)
k_pos = torch.arange(kv_len, device=device).unsqueeze(0)
return (k_pos <= q_pos) & (k_pos >= q_pos - sliding_window_size)
def init_forward_metadata(self, forward_batch: ForwardBatch): def init_forward_metadata(self, forward_batch: ForwardBatch):
"""Init the metadata for a forward pass.""" """Init the metadata for a forward pass."""
pass pass
@@ -44,6 +59,7 @@ class TorchNativeAttnBackend(AttentionBackend):
enable_gqa=False, enable_gqa=False,
causal=False, causal=False,
is_cross_attn=False, is_cross_attn=False,
sliding_window_size: Optional[int] = None,
): ):
"""Run the extend forward by using torch native sdpa op. """Run the extend forward by using torch native sdpa op.
@@ -114,14 +130,26 @@ class TorchNativeAttnBackend(AttentionBackend):
per_req_key = per_req_key.to(per_req_query.dtype) per_req_key = per_req_key.to(per_req_query.dtype)
per_req_value = per_req_value.to(per_req_query.dtype) per_req_value = per_req_value.to(per_req_query.dtype)
attn_mask = None
is_causal = causal
if sliding_window_size is not None and sliding_window_size > -1:
attn_mask = self._make_sliding_window_mask(
q_len=seq_len_kv,
kv_len=seq_len_kv,
sliding_window_size=sliding_window_size,
device=per_req_query.device,
)
is_causal = False
per_req_out_redudant = ( per_req_out_redudant = (
scaled_dot_product_attention( scaled_dot_product_attention(
per_req_query_redudant.unsqueeze(0), per_req_query_redudant.unsqueeze(0),
per_req_key.unsqueeze(0), per_req_key.unsqueeze(0),
per_req_value.unsqueeze(0), per_req_value.unsqueeze(0),
attn_mask=attn_mask,
enable_gqa=enable_gqa, enable_gqa=enable_gqa,
scale=scaling, scale=scaling,
is_causal=causal, is_causal=is_causal,
) )
.squeeze(0) .squeeze(0)
.movedim(query.dim() - 2, 0) .movedim(query.dim() - 2, 0)
@@ -144,6 +172,7 @@ class TorchNativeAttnBackend(AttentionBackend):
enable_gqa=False, enable_gqa=False,
causal=False, causal=False,
is_cross_attn=False, is_cross_attn=False,
sliding_window_size: Optional[int] = None,
): ):
"""Run the decode forward by using torch native sdpa op. """Run the decode forward by using torch native sdpa op.
@@ -202,14 +231,27 @@ class TorchNativeAttnBackend(AttentionBackend):
per_req_key = per_req_key.to(per_req_query.dtype) per_req_key = per_req_key.to(per_req_query.dtype)
per_req_value = per_req_value.to(per_req_query.dtype) per_req_value = per_req_value.to(per_req_query.dtype)
attn_mask = None
is_causal = causal
if sliding_window_size is not None and sliding_window_size > -1:
attn_mask = self._make_sliding_window_mask(
q_len=seq_len_q,
kv_len=seq_len_kv,
sliding_window_size=sliding_window_size,
device=per_req_query.device,
query_offset=seq_len_kv - seq_len_q,
)
is_causal = False
per_req_out = ( per_req_out = (
scaled_dot_product_attention( scaled_dot_product_attention(
per_req_query.unsqueeze(0), per_req_query.unsqueeze(0),
per_req_key.unsqueeze(0), per_req_key.unsqueeze(0),
per_req_value.unsqueeze(0), per_req_value.unsqueeze(0),
attn_mask=attn_mask,
enable_gqa=enable_gqa, enable_gqa=enable_gqa,
scale=scaling, scale=scaling,
is_causal=causal, is_causal=is_causal,
) )
.squeeze(0) .squeeze(0)
.movedim(query.dim() - 2, 0) .movedim(query.dim() - 2, 0)
@@ -265,6 +307,14 @@ class TorchNativeAttnBackend(AttentionBackend):
enable_gqa=use_gqa, enable_gqa=use_gqa,
causal=causal, causal=causal,
is_cross_attn=layer.is_cross_attention, is_cross_attn=layer.is_cross_attention,
sliding_window_size=(
layer.sliding_window_size
if causal
and not layer.is_cross_attention
and layer.sliding_window_size is not None
and layer.sliding_window_size > -1
else None
),
) )
return o return o
@@ -317,6 +367,13 @@ class TorchNativeAttnBackend(AttentionBackend):
enable_gqa=use_gqa, enable_gqa=use_gqa,
causal=False, causal=False,
is_cross_attn=layer.is_cross_attention, is_cross_attn=layer.is_cross_attention,
sliding_window_size=(
layer.sliding_window_size
if not layer.is_cross_attention
and layer.sliding_window_size is not None
and layer.sliding_window_size > -1
else None
),
) )
return o return o