perf: increase prefix attention tile budget 96MB→256MB

CCCL dispatch_transform.cuh spread_out_items_per_thread pattern:
reduce tile count = reduce Python loop iterations = faster prefill.

At 256K context with q_len=4096: old 96MB budget → 219 KV tokens/tile
→ ~1200 tiles per layer → 16 min per chunk. New 256MB budget →
~580 KV tokens/tile → ~450 tiles per layer → ~6 min per chunk.

BI-V100 has 32 GB HBM; 256 MB temporary tensor is safe.
This commit is contained in:
dylanyunlon
2026-08-07 04:44:18 +00:00
parent 025059d78e
commit e9eaad0592

View File

@@ -576,7 +576,9 @@ class PagedAttention:
# tile_sz=256 → 24 MB (safe)
# For decode (q_len=1): tile_sz=4096 → only 96 KB (always safe)
# ================================================================
_SMEM_BUDGET_BYTES = 96 * 1024 * 1024 # 96 MB score tensor budget
_SMEM_BUDGET_BYTES = 256 * 1024 * 1024 # 256 MB score tensor budget
# CCCL GridEvenShare: fewer tiles = fewer iterations = less overhead
# BI-V100 has 32 GB HBM per card; 256 MB temporary is safe.
batch_size = seq_lens_tensor.shape[0]
num_q_heads = query.shape[1]