[OPT] head_dim=256 Triton support — BLOCK=32 for Qwen3.6
CRITICAL DISCOVERY: Qwen3.6-35B-A3B uses head_dim=256 (not 128). text_cfg.head_dim=256, num_heads=24, num_kv_heads=4, GQA=6 This means ALL previous SMEM calculations were wrong: BLOCK=64 + head_dim=256: 64×256×2×2 = 64KB > 48KB → OVERFLOW BLOCK=64 + head_dim=128: 64×128×2×2 = 32KB ≤ 48KB → OK (but wrong model) Fix: head_dim-dependent BLOCK selection in prefix_prefill.py: head_dim ≤ 128: BLOCK=64, NUM_WARPS=4 (32KB SMEM) head_dim = 256: BLOCK=32, NUM_WARPS=4 (32KB SMEM) head_dim > 256: BLOCK=16, NUM_WARPS=2 (16KB SMEM) Also: _Q_CHUNK in _run_sdpa_fallback reduced 256→128 for head_dim=256 to avoid OOM on long sequences (256×100K×24×4=2.3GB vs 128×100K×24×4=1.2GB). Without this patch, Triton prefill CANNOT work for Qwen3.6. patch_enable_triton.py's try/fallback would always fall back to PyTorch.
This commit is contained in:
@@ -23,6 +23,11 @@ RUN python3 /workspace/qwen3_6_scripts/patch_triton_tuning.py
|
||||
# Triton Flash Attention is 10-50x faster than PyTorch for-loop fallback
|
||||
RUN python3 /workspace/qwen3_6_scripts/patch_enable_triton.py
|
||||
|
||||
# 5. head_dim=256 support: Qwen3.6 uses head_dim=256
|
||||
# BLOCK=64 overflows SMEM (64×256×2×2=64KB > 48KB)
|
||||
# → BLOCK=32 for head_dim=256 (32×256×2×2=32KB ≤ 48KB)
|
||||
RUN python3 /workspace/qwen3_6_scripts/patch_head256_triton.py
|
||||
|
||||
# 4. Raise decode threshold: compiled paged_attention_v1 up to 65536
|
||||
# instead of falling back to Python at 32768
|
||||
RUN python3 /workspace/qwen3_6_scripts/patch_vectorized_decode.py
|
||||
|
||||
182
qwen3_6_scripts/patch_head256_triton.py
Normal file
182
qwen3_6_scripts/patch_head256_triton.py
Normal file
@@ -0,0 +1,182 @@
|
||||
"""
|
||||
patch_head256_triton.py — Enable Triton prefill for head_dim=256
|
||||
=================================================================
|
||||
|
||||
Qwen3.6-35B-A3B uses head_dim=256 (confirmed: text_cfg.head_dim=256,
|
||||
num_heads=24, num_kv_heads=4, GQA ratio=6).
|
||||
|
||||
Current problem:
|
||||
prefix_prefill.py: BLOCK = 64 (set by patch_triton_tuning.py)
|
||||
SMEM needed: BLOCK_N × head_dim × sizeof(fp16) × 2 (K+V tiles)
|
||||
= 64 × 256 × 2 × 2 = 64KB > 48KB (BI-V100 SMEM limit)
|
||||
→ Triton kernel CANNOT launch at BLOCK=64 for head_dim=256.
|
||||
|
||||
xformers.py: _run_sdpa_fallback triggers for head_size > 128.
|
||||
This is a Python for-loop with Q-tiling — orders of magnitude slower.
|
||||
|
||||
Fix:
|
||||
1. In prefix_prefill.py launcher, use BLOCK based on head_dim:
|
||||
head_dim ≤ 128: BLOCK = 64, NUM_WARPS = 4 (32KB SMEM, fits)
|
||||
head_dim = 256: BLOCK = 32, NUM_WARPS = 4 (32KB SMEM, fits)
|
||||
head_dim > 256: BLOCK = 16, NUM_WARPS = 2 (16KB SMEM, fits)
|
||||
|
||||
2. Triton BLOCK=32 means more kernel launches per sequence but
|
||||
each launch uses only 32KB SMEM — well within 48KB limit.
|
||||
32×256×2×2 = 32KB ≤ 48KB ✓
|
||||
|
||||
3. Also optimize _run_sdpa_fallback _Q_CHUNK:
|
||||
Current: 256 (same for all head dims)
|
||||
For head_dim=256: memory = _Q_CHUNK × seq_len × H × 4 bytes (float32)
|
||||
At Q_CHUNK=256, seq_len=100K, H=24: 256×100K×24×4 ≈ 2.3GB
|
||||
Better: _Q_CHUNK=128 for head_dim=256 → 1.2GB (safer for OOM)
|
||||
|
||||
Deploy: python3 qwen3_6_scripts/patch_head256_triton.py
|
||||
Must run AFTER patch_ops.sh and patch_triton_tuning.py
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
PREFIX_PREFILL_PATHS = [
|
||||
"/usr/local/corex/lib/python3/dist-packages/vllm/attention/ops/prefix_prefill.py",
|
||||
"/usr/local/corex/lib64/python3/dist-packages/vllm/attention/ops/prefix_prefill.py",
|
||||
]
|
||||
|
||||
XFORMERS_PATHS = [
|
||||
"/usr/local/corex/lib/python3/dist-packages/vllm/attention/backends/xformers.py",
|
||||
"/usr/local/corex/lib64/python3/dist-packages/vllm/attention/backends/xformers.py",
|
||||
]
|
||||
|
||||
# --- Patch 1: prefix_prefill.py BLOCK selection based on head_dim ---
|
||||
|
||||
# The current patch_triton_tuning.py sets:
|
||||
# BLOCK = 64
|
||||
# NUM_WARPS = 4
|
||||
# We need to make BLOCK depend on head_dim:
|
||||
|
||||
OLD_BLOCK_SETTING = """\
|
||||
# BI-V100 optimization (patch_triton_tuning.py):
|
||||
# BLOCK=64: SMEM constrains BLOCK_N≤64 for head_dim=128
|
||||
# NUM_WARPS=4: fewer warps → more blocks/SM → better occupancy
|
||||
BLOCK = 64
|
||||
NUM_WARPS = 4"""
|
||||
|
||||
NEW_BLOCK_SETTING = """\
|
||||
# BI-V100: BLOCK must fit in 48KB SMEM.
|
||||
# SMEM = BLOCK_N × head_dim × sizeof(fp16) × 2 (K+V tiles)
|
||||
# head_dim=128: BLOCK=64 → 64×128×2×2 = 32KB ✓
|
||||
# head_dim=256: BLOCK=32 → 32×256×2×2 = 32KB ✓ (BLOCK=64 → 64KB overflow!)
|
||||
# head_dim>256: BLOCK=16 → fallback
|
||||
Lk = q.shape[-1]
|
||||
if Lk <= 128:
|
||||
BLOCK = 64
|
||||
NUM_WARPS = 4
|
||||
elif Lk <= 256:
|
||||
BLOCK = 32
|
||||
NUM_WARPS = 4
|
||||
else:
|
||||
BLOCK = 16
|
||||
NUM_WARPS = 2"""
|
||||
|
||||
# Alternative: if patch_triton_tuning.py hasn't run yet, patch the original
|
||||
OLD_BLOCK_ORIGINAL = """\
|
||||
BLOCK = 128 if current_platform.has_device_capability(80) else 64
|
||||
NUM_WARPS = 8"""
|
||||
|
||||
NEW_BLOCK_FROM_ORIGINAL = NEW_BLOCK_SETTING
|
||||
|
||||
|
||||
# --- Patch 2: _run_sdpa_fallback Q_CHUNK for head_dim=256 ---
|
||||
|
||||
OLD_Q_CHUNK = " _Q_CHUNK = 256"
|
||||
NEW_Q_CHUNK = """\
|
||||
# Adapt Q chunk size to head_dim to control memory:
|
||||
# head_dim=128: 256 × seq_len × H × 4B → manageable
|
||||
# head_dim=256: halve to 128 to avoid OOM on long sequences
|
||||
_Q_CHUNK = 128 if self.head_size > 128 else 256"""
|
||||
|
||||
|
||||
def patch_prefix_prefill():
|
||||
for path in PREFIX_PREFILL_PATHS:
|
||||
if not os.path.exists(path):
|
||||
continue
|
||||
with open(path, "r") as f:
|
||||
content = f.read()
|
||||
|
||||
changed = False
|
||||
if "head_dim=256: BLOCK=32" in content:
|
||||
print(f" [skip] {path}: head_dim-aware BLOCK already present")
|
||||
return True
|
||||
|
||||
if OLD_BLOCK_SETTING in content:
|
||||
content = content.replace(OLD_BLOCK_SETTING, NEW_BLOCK_SETTING, 1)
|
||||
changed = True
|
||||
print(f" [ok] Replaced fixed BLOCK=64 with head_dim-dependent selection")
|
||||
elif OLD_BLOCK_ORIGINAL in content:
|
||||
content = content.replace(OLD_BLOCK_ORIGINAL, NEW_BLOCK_FROM_ORIGINAL, 1)
|
||||
changed = True
|
||||
print(f" [ok] Replaced original BLOCK selection with head_dim-dependent version")
|
||||
else:
|
||||
print(f" [warn] Neither BLOCK anchor found in {path}")
|
||||
|
||||
# Also need to move Lk computation before BLOCK selection
|
||||
# Currently Lk is computed AFTER BLOCK is set (line ~750)
|
||||
# We need it before. Check if Lk is already available:
|
||||
if "Lk = q.shape[-1]" in content and "Lk, Lk, Lv" in content:
|
||||
# Lk is computed later — we duplicate the computation for BLOCK selection
|
||||
# This is safe because q.shape[-1] doesn't change
|
||||
print(f" [note] Lk computed early for BLOCK selection + later for kernel args")
|
||||
|
||||
if changed:
|
||||
with open(path, "w") as f:
|
||||
f.write(content)
|
||||
print(f" Written: {path}")
|
||||
return changed
|
||||
|
||||
print(" [error] prefix_prefill.py not found")
|
||||
return False
|
||||
|
||||
|
||||
def patch_xformers_sdpa():
|
||||
for path in XFORMERS_PATHS:
|
||||
if not os.path.exists(path):
|
||||
continue
|
||||
with open(path, "r") as f:
|
||||
content = f.read()
|
||||
|
||||
if "head_size > 128 else 256" in content:
|
||||
print(f" [skip] {path}: adaptive Q_CHUNK already present")
|
||||
return True
|
||||
|
||||
if OLD_Q_CHUNK in content:
|
||||
content = content.replace(OLD_Q_CHUNK, NEW_Q_CHUNK, 1)
|
||||
with open(path, "w") as f:
|
||||
f.write(content)
|
||||
print(f" [ok] {path}: _Q_CHUNK now adapts to head_dim")
|
||||
return True
|
||||
else:
|
||||
print(f" [warn] _Q_CHUNK anchor not found in {path}")
|
||||
|
||||
print(" [error] xformers.py not found")
|
||||
return False
|
||||
|
||||
|
||||
def main():
|
||||
print("=== patch_head256_triton: Enable Triton for head_dim=256 ===")
|
||||
print(f" Qwen3.6: head_dim=256, num_heads=24, num_kv_heads=4, GQA=6")
|
||||
print()
|
||||
|
||||
print("--- Patch 1: prefix_prefill.py BLOCK selection ---")
|
||||
patch_prefix_prefill()
|
||||
|
||||
print("\n--- Patch 2: xformers Q_CHUNK for head_dim=256 ---")
|
||||
patch_xformers_sdpa()
|
||||
|
||||
print("\nSMEM budget at BLOCK=32, head_dim=256:")
|
||||
print(f" K tile: 32 × 256 × 2B = 16KB")
|
||||
print(f" V tile: 32 × 256 × 2B = 16KB")
|
||||
print(f" Total: 32KB ≤ 48KB ✓")
|
||||
print("\nDone.")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user