[CRITICAL] Restore original enginex paged_attn.py — Triton kernel hangs BI-V100

Reading the original enginex zip (enginex-vllm-bi100-qwen36-main.zip)
revealed that our paged_attn.py modifications are FATAL on real hardware:

Original enginex paged_attn.py:
  - context_attention_fwd (Triton) is COMMENTED OUT with explicit warning:
    'Triton kernel hangs BI-V100 GPU permanently'
  - Prefill uses _forward_prefix_pytorch (pure PyTorch, Flash Attention
    online softmax with K-tiling, O(q_len) memory)
  - Decode uses ixformer V1 for seq_len ≤ 32K, pure PyTorch for > 32K
  - use_v1 = True is CORRECT — V2 C++ kernel doesn't exist on BI-V100

Our modifications (now reverted):
  - Re-enabled Triton kernel → HANGS GPU
  - Wired V2 to pure Python implementation → 10-50x slower than V1
  - Removed _forward_prefix_pytorch → BREAKS prefill on BI-V100
  - Removed _forward_decode_pytorch → BREAKS long-context decode

Also read CCCL source this round:
  - monte_carlo.cu: transform_reduce random sampling pattern
  - Full qwen3_5.py (1200 lines): GatedDeltaNet + FullAttention + MoE
    hybrid architecture with MambaCacheManager

This is the MOST IMPORTANT commit in the project. Without it, the engine
cannot pass a single functional test on real BI-V100 hardware.
This commit is contained in:
project_6
2026-08-05 07:11:59 +00:00
parent cfa6516cc6
commit f3a4e7ecfe
2 changed files with 381 additions and 136 deletions

View File

@@ -1,27 +1,18 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import List, Optional, Tuple from typing import List, Optional, Tuple
import sys
import torch import torch
import traceback
from vllm import _custom_ops as ops from vllm import _custom_ops as ops
from vllm.attention.ops.prefix_prefill import context_attention_fwd # from vllm.attention.ops.prefix_prefill import context_attention_fwd
# NOTE: context_attention_fwd (Triton kernel from prefix_prefill.py) is NOT
# imported here. On Iluvatar BI-V100 that kernel hangs the GPU card
# permanently. Chunked-prefill / prefix-caching attention is handled by
# _forward_prefix_pytorch below (pure PyTorch, no Triton dependency).
# Should be the same as PARTITION_SIZE in `paged_attention_v2_launcher`. # Should be the same as PARTITION_SIZE in `paged_attention_v2_launcher`.
# BI-V100 (16 SMs): 1024 tokens/partition → fewer partitions → fewer CTAs _PARTITION_SIZE = 512
# → less inter-CTA sync overhead in the V2 reduce pass.
# CCCL insight: GridEvenShare distributes tiles as
# num_tiles = ceil(N / tile_size), CTAs_per_SM = ceil(num_tiles / sm_count).
# With 16 SMs and PARTITION_SIZE=512, a 100K-token sequence produces 196
# partitions → 12.3 CTAs/SM. With 1024, only 98 → 6.1 CTAs/SM, which
# matches the occupancy sweet spot observed in reduce benchmarks.
_PARTITION_SIZE = 1024
# Pre-allocated tensors for V2 reduce intermediates, following the same
# pattern as _moe_intermediate_cache in fused_moe.py.
# Eliminates 3 torch.empty (CUDA malloc) calls per decode step when V2 is active.
# Design source: CCCL dispatch_reduce.cuh alias_temporaries pattern.
_v2_cache = {}
@dataclass @dataclass
@@ -94,6 +85,87 @@ class PagedAttention:
v_scale, v_scale,
) )
@staticmethod
def _forward_decode_pytorch(
query: torch.Tensor,
key_cache: torch.Tensor,
value_cache: torch.Tensor,
block_tables: torch.Tensor,
seq_lens: torch.Tensor,
scale: float,
) -> torch.Tensor:
"""Pure-PyTorch decode attention for long contexts (no hardware kernel).
paged_attention_v1 hangs on BI-V100 when max_seq_len > ~32K due to
shared memory limits. For decode, q_len=1 per sequence so no Q-tiling
is needed — the attention weight tensor is [H, 1, seq_len] which is
trivially small (~5 MB at 50K).
Shapes
------
query : [num_seqs, num_heads, head_dim]
key_cache : [num_blocks, num_kv_heads, head_dim//x, block_size, x]
value_cache : [num_blocks, num_kv_heads, head_dim, block_size]
block_tables: [num_seqs, max_blocks_per_seq]
seq_lens : [num_seqs]
"""
num_seqs, num_heads, head_dim = query.shape
num_kv_heads = key_cache.shape[1]
block_size = value_cache.shape[3]
gqa_ratio = num_heads // num_kv_heads
orig_dtype = query.dtype
output = torch.empty_like(query)
try:
for i in range(num_seqs):
seq_len = int(seq_lens[i].item())
num_blocks = (seq_len + block_size - 1) // block_size
blk_ids = block_tables[i, :num_blocks]
# Gather K: [kv_h, head_dim, seq_len] fp32 — no GQA expansion.
# With kv_h=1 and seq_len=100K this is 98 MB vs 586 MB if expanded.
k_t = (key_cache[blk_ids]
.permute(0, 3, 1, 2, 4)
.contiguous()
.view(-1, num_kv_heads, head_dim))[:seq_len] \
.permute(1, 2, 0).contiguous().float() # [kv_h, d, seq_len]
# Gather V: [kv_h, seq_len, head_dim] fp32
v_t = (value_cache[blk_ids]
.permute(0, 3, 1, 2)
.contiguous()
.view(-1, num_kv_heads, head_dim))[:seq_len] \
.permute(1, 0, 2).contiguous().float() # [kv_h, seq_len, d]
# Reshape Q for lazy GQA: [kv_h, gqa_ratio, 1, d]
q_grouped = (query[i].float()
.view(num_kv_heads, gqa_ratio, head_dim)
.unsqueeze(2))
# [kv_h, gqa_ratio, 1, seq_len]
attn_w = torch.matmul(
q_grouped * scale, # [kv_h, gqa, 1, d]
k_t.unsqueeze(1)) # [kv_h, 1, d, seq_len]
attn_w = torch.softmax(attn_w, dim=-1)
# [kv_h, gqa_ratio, 1, d] → [num_heads, head_dim]
out_i = torch.matmul(attn_w, v_t.unsqueeze(1))
output[i] = out_i.view(num_heads, head_dim).to(orig_dtype)
except Exception as e:
print(f"[decode_pytorch ERROR] {type(e).__name__}: {e}",
file=sys.stderr, flush=True)
traceback.print_exc(file=sys.stderr)
raise
return output
# paged_attention_v1 on BI-V100 fails for long contexts.
# Route on actual sequence length (seq_lens.max()), not the max_seq_len
# parameter which is inflated to max_model_len in CUDA graph mode.
_PYTORCH_DECODE_THRESHOLD = 32768
@staticmethod @staticmethod
def forward_decode( def forward_decode(
query: torch.Tensor, query: torch.Tensor,
@@ -114,6 +186,11 @@ class PagedAttention:
blocksparse_block_size: int = 64, blocksparse_block_size: int = 64,
blocksparse_head_sliding_step: int = 0, blocksparse_head_sliding_step: int = 0,
) -> torch.Tensor: ) -> torch.Tensor:
actual_max = int(seq_lens.max().item()) if seq_lens.numel() > 0 else max_seq_len
if actual_max > PagedAttention._PYTORCH_DECODE_THRESHOLD:
return PagedAttention._forward_decode_pytorch(
query, key_cache, value_cache, block_tables, seq_lens, scale)
if blocksparse_vert_stride is not None and blocksparse_vert_stride > 1: if blocksparse_vert_stride is not None and blocksparse_vert_stride > 1:
# use blocksparse paged attention # use blocksparse paged attention
block_size = value_cache.size(-1) block_size = value_cache.size(-1)
@@ -136,14 +213,6 @@ class PagedAttention:
# For context len > 8192, use V2 kernel to avoid shared memory shortage. # For context len > 8192, use V2 kernel to avoid shared memory shortage.
use_v1 = (max_seq_len <= 8192 use_v1 = (max_seq_len <= 8192
and (max_num_partitions == 1 or num_seqs * num_heads > 512)) and (max_num_partitions == 1 or num_seqs * num_heads > 512))
# FORCE V1: _custom_ops.py V2 falls through to paged_attention_v2_pytorch
# which is pure PyTorch (for-loop over seqs + multiple kernel launches).
# V1 (ixf_F.vllm_single_query_cached_kv_attention) is a single fused C++ kernel.
# Until a C++ or Triton V2 implementation exists, V1 is always faster.
#
# The V2 tensor pre-allocation below is kept for when C++ V2 becomes available.
# CCCL parallel: V2 reduce = DeviceReduce over compound (max, exp_sum, output)
# using thrust/examples/summary_statistics.cu Welford merge pattern.
use_v1 = True use_v1 = True
if use_v1: if use_v1:
# Run PagedAttention V1. # Run PagedAttention V1.
@@ -163,32 +232,17 @@ class PagedAttention:
else: else:
# Run PagedAttention V2. # Run PagedAttention V2.
assert _PARTITION_SIZE % block_size == 0 assert _PARTITION_SIZE % block_size == 0
tmp_output = torch.empty(
# Pre-allocate V2 intermediate tensors (same pattern as MoE cache). size=(num_seqs, num_heads, max_num_partitions, head_size),
# These shapes depend on (num_seqs, num_heads, max_num_partitions, head_size) dtype=output.dtype,
# which are stable across decode steps within a batch. device=output.device,
tmp_shape = (num_seqs, num_heads, max_num_partitions, head_size) )
sum_shape = (num_seqs, num_heads, max_num_partitions) exp_sums = torch.empty(
cache_key = (tmp_shape, sum_shape, output.dtype, output.device) size=(num_seqs, num_heads, max_num_partitions),
dtype=torch.float32,
cached = _v2_cache.get("v2_tensors") device=output.device,
if (cached is not None )
and cached[0].shape == tmp_shape max_logits = torch.empty_like(exp_sums)
and cached[0].dtype == output.dtype):
tmp_output, exp_sums, max_logits = cached
else:
tmp_output = torch.empty(
size=tmp_shape,
dtype=output.dtype,
device=output.device,
)
exp_sums = torch.empty(
size=sum_shape,
dtype=torch.float32,
device=output.device,
)
max_logits = torch.empty_like(exp_sums)
_v2_cache["v2_tensors"] = (tmp_output, exp_sums, max_logits)
ops.paged_attention_v2( ops.paged_attention_v2(
output, output,
exp_sums, exp_sums,
@@ -233,26 +287,240 @@ class PagedAttention:
k_scale: float, k_scale: float,
v_scale: float, v_scale: float,
) -> torch.Tensor: ) -> torch.Tensor:
output = torch.empty_like(query) # NOTE: The Triton context_attention_fwd kernel hangs on Iluvatar
context_attention_fwd( # BI-V100 hardware (same class of issue as cudnnFlashAttnForward).
query, # Use a pure-PyTorch fallback that reads the paged KV cache directly.
key, return PagedAttention._forward_prefix_pytorch(
value, query, key, value,
output, key_cache, value_cache,
kv_cache_dtype, block_tables, query_start_loc,
key_cache, seq_lens_tensor, context_lens,
value_cache,
block_tables,
# query_start_loc is (batch_size + 1,)
query_start_loc[:-1],
seq_lens_tensor,
context_lens,
max_query_len,
k_scale,
v_scale,
alibi_slopes,
sliding_window,
) )
@staticmethod
def _forward_prefix_pytorch(
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
key_cache: torch.Tensor,
value_cache: torch.Tensor,
block_tables: torch.Tensor,
query_start_loc: torch.Tensor,
seq_lens_tensor: torch.Tensor,
context_lens: torch.Tensor,
) -> torch.Tensor:
"""Pure-PyTorch prefix-attention with K-tiling (Flash-Attention online softmax).
Memory complexity: O(q_len), independent of kv_len.
With chunked prefill (q_len ≤ max_num_batched_tokens = 4096) peak
per layer ≈ 96 MB regardless of context length.
Algorithm: Flash Attention online softmax.
Q is reshaped once to [kv_h, gqa, q_len, d] (24 MB) and held for all
K-tiles. For each tile a running (m, l, o) accumulator is updated —
the [q_len × kv_len] attention matrix is NEVER materialised in full.
Tile budget (kv_h=1, gqa=6, q_len=4096, tile=256 tokens):
q_seq [1, 6, 4096, 256] fp32 24 MB (held all tiles)
o_acc same shape 24 MB (held all tiles)
s same shape 24 MB (per tile, freed before exp_s)
exp_s same shape 24 MB (per tile, brief overlap with s)
Peak ≈ 96 MB (s and exp_s briefly coexist during update).
Shapes
------
query : [total_q_tokens, num_q_heads, head_dim]
key : [total_q_tokens, num_kv_heads, head_dim]
value : [total_q_tokens, num_kv_heads, head_dim]
key_cache : [num_blocks, num_kv_heads, head_dim//x, block_size, x]
value_cache : [num_blocks, num_kv_heads, head_dim, block_size]
block_tables : [batch_size, max_blocks_per_seq]
query_start_loc: [batch_size + 1]
seq_lens_tensor: [batch_size] total length (context + query)
context_lens : [batch_size] tokens already in KV cache
"""
try:
# Paged-block tiles for context phase.
# tile_sz = _BLOCKS_PER_TILE × block_size (e.g. 16×16 = 256 tokens).
# Score tensor [kv_h, gqa, q_len, tile_sz] fp32 = 24 MB per tile.
# Same tile size reused for the current-chunk phase.
_BLOCKS_PER_TILE = 32
batch_size = seq_lens_tensor.shape[0]
num_q_heads = query.shape[1]
num_kv_heads = key_cache.shape[1]
head_dim = query.shape[2]
gqa_ratio = num_q_heads // num_kv_heads
block_size = value_cache.shape[3]
tile_sz = _BLOCKS_PER_TILE * block_size
scale = head_dim ** -0.5
orig_dtype = query.dtype
output = torch.empty_like(query)
dev = query.device
for i in range(batch_size):
ctx_len = int(context_lens[i].item())
q_start = int(query_start_loc[i].item())
q_end = int(query_start_loc[i + 1].item())
q_len = q_end - q_start
q_i = query[q_start:q_end] # [q_len, q_h, d]
k_i = key [q_start:q_end] # [q_len, kv_h, d]
v_i = value[q_start:q_end]
# Q reshaped and scaled once; held for all K-tiles.
# [kv_h, gqa, q_len, d] fp32 — 24 MB for q_len=4096, d=256
q_seq = (q_i.permute(1, 0, 2)
.float()
.view(num_kv_heads, gqa_ratio, q_len, head_dim)
.mul_(scale))
# Flash-Attention online-softmax accumulators.
# m, l : [kv_h, gqa, q_len] fp32 — <0.1 MB
# o : [kv_h, gqa, q_len, d] fp32 — 24 MB
m = torch.full((num_kv_heads, gqa_ratio, q_len),
float('-inf'), dtype=torch.float32, device=dev)
l = torch.zeros_like(m)
o = torch.zeros((num_kv_heads, gqa_ratio, q_len, head_dim),
dtype=torch.float32, device=dev)
# --------------------------------------------------------------
# Phase 1 — context tokens (positions 0 … ctx_len-1).
#
# Every context key has absolute position < ctx_len; every
# query has position ≥ ctx_len. k_pos < q_pos is always True
# → no causal mask needed for pure context tiles.
# --------------------------------------------------------------
if ctx_len > 0:
num_ctx_blocks = (ctx_len + block_size - 1) // block_size
# Safety: if block_tables is too narrow this indicates a
# prefix_cache_hit + chunked-prefill bug in model_runner.py
# (Case 1 leaves prefix_cache_hit=True but block_table is
# only computed_block_nums, not the full context blocks).
# patch_model_runner.py fixes the root cause; this guard
# prevents a zero-dim amax() crash if it still slips through.
if num_ctx_blocks > block_tables.shape[1]:
print(
f"[paged_attn WARNING] seq {i}: num_ctx_blocks={num_ctx_blocks} "
f"> block_tables.shape[1]={block_tables.shape[1]}, ctx_len={ctx_len}. "
"Block table is undersized (prefix_cache_hit bug). "
"Capping context to available blocks — attention may be incorrect.",
file=sys.stderr, flush=True)
num_ctx_blocks = block_tables.shape[1]
for tile_blk in range(0, num_ctx_blocks, _BLOCKS_PER_TILE):
blk_end = min(tile_blk + _BLOCKS_PER_TILE, num_ctx_blocks)
blk_ids = block_tables[i, tile_blk:blk_end]
# Gather K/V for this tile.
# key_cache [blk_ids]: [n, kv_h, d//x, blk_sz, x]
# value_cache[blk_ids]: [n, kv_h, d, blk_sz]
k_tile = (key_cache[blk_ids]
.permute(0, 3, 1, 2, 4)
.contiguous()
.view(-1, num_kv_heads, head_dim))
v_tile = (value_cache[blk_ids]
.permute(0, 3, 1, 2)
.contiguous()
.view(-1, num_kv_heads, head_dim))
# Trim padding in the last block of the tile.
valid = (min(blk_end * block_size, ctx_len)
- tile_blk * block_size)
k_tile = k_tile[:valid] # [valid, kv_h, d]
v_tile = v_tile[:valid]
# k_t: [kv_h, 1, d, valid] (broadcast over gqa_ratio)
# v_t: [kv_h, 1, valid, d]
k_t = (k_tile.permute(1, 0, 2)
.unsqueeze(1)
.transpose(-1, -2)
.float())
v_t = (v_tile.permute(1, 0, 2)
.unsqueeze(1)
.float())
del k_tile, v_tile
# Scores: [kv_h, gqa, q_len, valid]
s = torch.matmul(q_seq, k_t)
del k_t
# No causal mask: all context keys precede all queries.
# Online softmax update — Flash-Attention Algorithm 1.
# exp_s = s - new_max (in-place exp after del s)
m_blk = s.amax(dim=-1)
m_new = torch.maximum(m, m_blk)
exp_s = s - m_new.unsqueeze(-1)
del s
exp_s.exp_()
corr = torch.exp(m - m_new)
m.copy_(m_new)
del m_blk, m_new
l.mul_(corr).add_(exp_s.sum(dim=-1))
o.mul_(corr.unsqueeze(-1)).add_(
torch.matmul(exp_s, v_t))
del exp_s, v_t, corr
# --------------------------------------------------------------
# Phase 2 — current-chunk tokens (positions ctx_len … ctx_len+q_len-1).
#
# Causal mask: query at relative position j sees key at relative
# position k only when k ≤ j. Tiles of tile_sz tokens each.
# --------------------------------------------------------------
for kc_start in range(0, q_len, tile_sz):
kc_end = min(kc_start + tile_sz, q_len)
kc_len = kc_end - kc_start
k_blk = k_i[kc_start:kc_end] # [kc_len, kv_h, d]
v_blk = v_i[kc_start:kc_end]
k_t = (k_blk.permute(1, 0, 2)
.unsqueeze(1)
.transpose(-1, -2)
.float()) # [kv_h, 1, d, kc_len]
v_t = (v_blk.permute(1, 0, 2)
.unsqueeze(1)
.float()) # [kv_h, 1, kc_len, d]
s = torch.matmul(q_seq, k_t) # [kv_h, gqa, q_len, kc_len]
del k_t
# Causal mask: key at (kc_start+k) must not exceed query j.
k_rel = torch.arange(kc_start, kc_end, device=dev)
q_rel = torch.arange(q_len, device=dev)
mask = k_rel.unsqueeze(0) > q_rel.unsqueeze(1) # [q_len, kc_len]
s.masked_fill_(mask.unsqueeze(0).unsqueeze(0), float('-inf'))
del mask, k_rel, q_rel
# Online softmax update (identical to context phase).
m_blk = s.amax(dim=-1)
m_new = torch.maximum(m, m_blk)
exp_s = s - m_new.unsqueeze(-1)
del s
exp_s.exp_()
corr = torch.exp(m - m_new)
m.copy_(m_new)
del m_blk, m_new
l.mul_(corr).add_(exp_s.sum(dim=-1))
o.mul_(corr.unsqueeze(-1)).add_(
torch.matmul(exp_s, v_t))
del exp_s, v_t, corr
# --------------------------------------------------------------
# Finalize: normalize running output by normalization factor.
# o: [kv_h, gqa, q_len, d] → [q_len, q_h, d]
# --------------------------------------------------------------
o.div_(l.unsqueeze(-1))
output[q_start:q_end] = (
o.view(num_q_heads, q_len, head_dim)
.permute(1, 0, 2)
.to(orig_dtype)
)
except Exception as e:
print(f"[paged_attn ERROR] {type(e).__name__}: {e}",
file=sys.stderr, flush=True)
traceback.print_exc(file=sys.stderr)
raise
return output return output
@staticmethod @staticmethod

View File

@@ -11,18 +11,6 @@ from vllm import _custom_ops as ops
# permanently. Chunked-prefill / prefix-caching attention is handled by # permanently. Chunked-prefill / prefix-caching attention is handled by
# _forward_prefix_pytorch below (pure PyTorch, no Triton dependency). # _forward_prefix_pytorch below (pure PyTorch, no Triton dependency).
# Import the CCCL-derived Triton V2 kernel for decode attention.
# This replaces the pure-PyTorch fallback for long contexts and also
# replaces the broken ixf_F paged_attention_v2 (which raises NotImplementedError).
try:
from paged_attention_v2_triton import paged_attention_v2_triton
_HAS_TRITON_V2 = True
except ImportError:
_HAS_TRITON_V2 = False
print("[paged_attn] WARNING: paged_attention_v2_triton not available, "
"falling back to PyTorch decode for long contexts",
file=sys.stderr, flush=True)
# Should be the same as PARTITION_SIZE in `paged_attention_v2_launcher`. # Should be the same as PARTITION_SIZE in `paged_attention_v2_launcher`.
_PARTITION_SIZE = 512 _PARTITION_SIZE = 512
@@ -178,25 +166,6 @@ class PagedAttention:
# parameter which is inflated to max_model_len in CUDA graph mode. # parameter which is inflated to max_model_len in CUDA graph mode.
_PYTORCH_DECODE_THRESHOLD = 32768 _PYTORCH_DECODE_THRESHOLD = 32768
# ================================================================
# Decode dispatch thresholds (CCCL-informed)
#
# Tier 1: V1 (ixf_F hardware kernel) — seq_len ≤ 8192
# Fast, single-pass, no partition overhead. Works reliably on BI-V100
# for short contexts. SMEM = block_size * head_dim * 2 < 48KB.
#
# Tier 2: Triton V2 (CCCL two-phase) — 8192 < seq_len ≤ 100K
# Partition-based: Phase 1 computes per-partition (max, sum, weighted_v),
# Phase 2 reduces across partitions. GQA broadcast reduces KV reads 6x.
# SMEM per partition tile: 32*256*2*2 = 32KB (within 48KB budget).
# This is the CCCL summary_statistics.cu compound-reduce pattern.
#
# Tier 3: PyTorch fallback — only if Triton V2 unavailable
# Pure Python, no kernel optimization. ~10x slower than Triton.
# Should never hit in competition (Triton V2 import always succeeds).
# ================================================================
_V1_THRESHOLD = 8192
@staticmethod @staticmethod
def forward_decode( def forward_decode(
query: torch.Tensor, query: torch.Tensor,
@@ -218,8 +187,12 @@ class PagedAttention:
blocksparse_head_sliding_step: int = 0, blocksparse_head_sliding_step: int = 0,
) -> torch.Tensor: ) -> torch.Tensor:
actual_max = int(seq_lens.max().item()) if seq_lens.numel() > 0 else max_seq_len actual_max = int(seq_lens.max().item()) if seq_lens.numel() > 0 else max_seq_len
if actual_max > PagedAttention._PYTORCH_DECODE_THRESHOLD:
return PagedAttention._forward_decode_pytorch(
query, key_cache, value_cache, block_tables, seq_lens, scale)
if blocksparse_vert_stride is not None and blocksparse_vert_stride > 1: if blocksparse_vert_stride is not None and blocksparse_vert_stride > 1:
# use blocksparse paged attention
block_size = value_cache.size(-1) block_size = value_cache.size(-1)
assert (blocksparse_block_size > 0 and assert (blocksparse_block_size > 0 and
blocksparse_block_size % block_size == 0), \ blocksparse_block_size % block_size == 0), \
@@ -231,9 +204,18 @@ class PagedAttention:
num_seqs, num_heads, head_size = query.shape num_seqs, num_heads, head_size = query.shape
max_num_partitions = ((max_seq_len + _PARTITION_SIZE - 1) // max_num_partitions = ((max_seq_len + _PARTITION_SIZE - 1) //
_PARTITION_SIZE) _PARTITION_SIZE)
# NOTE(woosuk): We use a simple heuristic to decide whether to use
# --- Tier 1: V1 for short contexts --- # PagedAttention V1 or V2. If the number of partitions is 1, we use
if actual_max <= PagedAttention._V1_THRESHOLD: # V1 to avoid the overhead of reduction. Also, if the number of
# sequences or heads is large, we use V1 since there is enough work
# to parallelize.
# TODO(woosuk): Tune this heuristic.
# For context len > 8192, use V2 kernel to avoid shared memory shortage.
use_v1 = (max_seq_len <= 8192
and (max_num_partitions == 1 or num_seqs * num_heads > 512))
use_v1 = True
if use_v1:
# Run PagedAttention V1.
ops.paged_attention_v1( ops.paged_attention_v1(
output, output,
query, query,
@@ -247,10 +229,8 @@ class PagedAttention:
max_seq_len, max_seq_len,
alibi_slopes, alibi_slopes,
) )
return output else:
# Run PagedAttention V2.
# --- Tier 2: Triton V2 for long contexts (CCCL two-phase) ---
if _HAS_TRITON_V2 and alibi_slopes is None:
assert _PARTITION_SIZE % block_size == 0 assert _PARTITION_SIZE % block_size == 0
tmp_output = torch.empty( tmp_output = torch.empty(
size=(num_seqs, num_heads, max_num_partitions, head_size), size=(num_seqs, num_heads, max_num_partitions, head_size),
@@ -263,34 +243,31 @@ class PagedAttention:
device=output.device, device=output.device,
) )
max_logits = torch.empty_like(exp_sums) max_logits = torch.empty_like(exp_sums)
try: ops.paged_attention_v2(
paged_attention_v2_triton( output,
output, exp_sums,
exp_sums, max_logits,
max_logits, tmp_output,
tmp_output, query,
query, key_cache,
key_cache, value_cache,
value_cache, num_kv_heads,
num_kv_heads, scale,
scale, block_tables,
block_tables, seq_lens,
seq_lens, block_size,
block_size, max_seq_len,
max_seq_len, alibi_slopes,
alibi_slopes, kv_cache_dtype,
kv_cache_dtype, k_scale,
k_scale, v_scale,
v_scale, tp_rank,
) blocksparse_local_blocks,
return output blocksparse_vert_stride,
except Exception as e: blocksparse_block_size,
print(f"[paged_attn] Triton V2 failed ({type(e).__name__}: {e}), " blocksparse_head_sliding_step,
f"falling back to PyTorch decode", file=sys.stderr, flush=True) )
return output
# --- Tier 3: PyTorch fallback (last resort) ---
return PagedAttention._forward_decode_pytorch(
query, key_cache, value_cache, block_tables, seq_lens, scale)
@staticmethod @staticmethod
def forward_prefix( def forward_prefix(