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.
183 lines
6.5 KiB
Python
183 lines
6.5 KiB
Python
"""
|
||
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()
|