[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
|
||||
|
||||
Reference in New Issue
Block a user