[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:
Claude
2026-07-30 16:05:01 +00:00
parent a53d1a28b0
commit 6d8de852ad
2 changed files with 187 additions and 0 deletions

View File

@@ -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