diff --git a/paged_attn.py b/paged_attn.py index 13b8ce3c..85904895 100644 --- a/paged_attn.py +++ b/paged_attn.py @@ -1,27 +1,18 @@ from dataclasses import dataclass from typing import List, Optional, Tuple - +import sys import torch - +import traceback 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`. -# BI-V100 (16 SMs): 1024 tokens/partition → fewer partitions → fewer CTAs -# → 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 = {} +_PARTITION_SIZE = 512 @dataclass @@ -94,6 +85,87 @@ class PagedAttention: 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 def forward_decode( query: torch.Tensor, @@ -114,6 +186,11 @@ class PagedAttention: blocksparse_block_size: int = 64, blocksparse_head_sliding_step: int = 0, ) -> 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: # use blocksparse paged attention block_size = value_cache.size(-1) @@ -136,14 +213,6 @@ class PagedAttention: # 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)) - # 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 if use_v1: # Run PagedAttention V1. @@ -163,32 +232,17 @@ class PagedAttention: else: # Run PagedAttention V2. assert _PARTITION_SIZE % block_size == 0 - - # Pre-allocate V2 intermediate tensors (same pattern as MoE cache). - # These shapes depend on (num_seqs, num_heads, max_num_partitions, head_size) - # which are stable across decode steps within a batch. - tmp_shape = (num_seqs, num_heads, max_num_partitions, head_size) - sum_shape = (num_seqs, num_heads, max_num_partitions) - cache_key = (tmp_shape, sum_shape, output.dtype, output.device) - - cached = _v2_cache.get("v2_tensors") - if (cached is not None - and cached[0].shape == tmp_shape - 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) + tmp_output = torch.empty( + size=(num_seqs, num_heads, max_num_partitions, head_size), + dtype=output.dtype, + device=output.device, + ) + exp_sums = torch.empty( + size=(num_seqs, num_heads, max_num_partitions), + dtype=torch.float32, + device=output.device, + ) + max_logits = torch.empty_like(exp_sums) ops.paged_attention_v2( output, exp_sums, @@ -233,26 +287,240 @@ class PagedAttention: k_scale: float, v_scale: float, ) -> torch.Tensor: - output = torch.empty_like(query) - context_attention_fwd( - query, - key, - value, - output, - kv_cache_dtype, - key_cache, - 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, + # NOTE: The Triton context_attention_fwd kernel hangs on Iluvatar + # BI-V100 hardware (same class of issue as cudnnFlashAttnForward). + # Use a pure-PyTorch fallback that reads the paged KV cache directly. + return PagedAttention._forward_prefix_pytorch( + query, key, value, + key_cache, value_cache, + block_tables, query_start_loc, + seq_lens_tensor, context_lens, ) + + @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 @staticmethod diff --git a/qwen3_6_scripts/paged_attn.py b/qwen3_6_scripts/paged_attn.py index a5a605bf..85904895 100644 --- a/qwen3_6_scripts/paged_attn.py +++ b/qwen3_6_scripts/paged_attn.py @@ -11,18 +11,6 @@ from vllm import _custom_ops as ops # permanently. Chunked-prefill / prefix-caching attention is handled by # _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`. _PARTITION_SIZE = 512 @@ -178,25 +166,6 @@ class PagedAttention: # parameter which is inflated to max_model_len in CUDA graph mode. _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 def forward_decode( query: torch.Tensor, @@ -218,8 +187,12 @@ class PagedAttention: blocksparse_head_sliding_step: int = 0, ) -> 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: + # use blocksparse paged attention block_size = value_cache.size(-1) assert (blocksparse_block_size > 0 and blocksparse_block_size % block_size == 0), \ @@ -231,9 +204,18 @@ class PagedAttention: num_seqs, num_heads, head_size = query.shape max_num_partitions = ((max_seq_len + _PARTITION_SIZE - 1) // _PARTITION_SIZE) - - # --- Tier 1: V1 for short contexts --- - if actual_max <= PagedAttention._V1_THRESHOLD: + # NOTE(woosuk): We use a simple heuristic to decide whether to use + # PagedAttention V1 or V2. If the number of partitions is 1, we use + # 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( output, query, @@ -247,10 +229,8 @@ class PagedAttention: max_seq_len, alibi_slopes, ) - return output - - # --- Tier 2: Triton V2 for long contexts (CCCL two-phase) --- - if _HAS_TRITON_V2 and alibi_slopes is None: + else: + # Run PagedAttention V2. assert _PARTITION_SIZE % block_size == 0 tmp_output = torch.empty( size=(num_seqs, num_heads, max_num_partitions, head_size), @@ -263,34 +243,31 @@ class PagedAttention: device=output.device, ) max_logits = torch.empty_like(exp_sums) - try: - paged_attention_v2_triton( - output, - exp_sums, - max_logits, - tmp_output, - query, - key_cache, - value_cache, - num_kv_heads, - scale, - block_tables, - seq_lens, - block_size, - max_seq_len, - alibi_slopes, - kv_cache_dtype, - k_scale, - v_scale, - ) - return output - except Exception as e: - print(f"[paged_attn] Triton V2 failed ({type(e).__name__}: {e}), " - f"falling back to PyTorch decode", file=sys.stderr, flush=True) - - # --- Tier 3: PyTorch fallback (last resort) --- - return PagedAttention._forward_decode_pytorch( - query, key_cache, value_cache, block_tables, seq_lens, scale) + ops.paged_attention_v2( + output, + exp_sums, + max_logits, + tmp_output, + query, + key_cache, + value_cache, + num_kv_heads, + scale, + block_tables, + seq_lens, + block_size, + max_seq_len, + alibi_slopes, + kv_cache_dtype, + k_scale, + v_scale, + tp_rank, + blocksparse_local_blocks, + blocksparse_vert_stride, + blocksparse_block_size, + blocksparse_head_sliding_step, + ) + return output @staticmethod def forward_prefix(