Files
project_6/qwen3_6_scripts/patch_triton_tuning.py
Claude 4463e9ccee [OPT] BI-V100 Triton kernel tuning + computility-run.yaml optimization
After reading the full baseline (enginex-vllm-bi100-qwen36-main.zip):

KEY DISCOVERY: The competition optimization surface is Python/Triton,
not C++ CUDA. There is no csrc/ directory. All CUDA kernels are
precompiled in vllm._C and ixformer .so files. The muh C++ headers
have no injection point in this competition framework.

What CAN be optimized:

1. Triton kernel parameters (prefix_prefill.py):
   - BLOCK: stays at 64 (correct — BLOCK_N=128 overflows 48KB SMEM
     at head_dim=128: 128×128×2×2=64KB > 48KB)
   - NUM_WARPS: 8 → 4 (derived from occupancy analysis:
     at 8 warps + 32KB SMEM/block, only 1 block fits per SM;
     at 4 warps, potentially 2 blocks per SM = 2× occupancy;
     BI-V100 is bandwidth-limited (900GB/s), so more blocks
     hiding bandwidth latency matters more than more warps
     hiding instruction latency)

2. computility-run.yaml:
   - max-num-batched-tokens: 8192 → 16384 (larger prefill chunks
     reduce kernel launch overhead; with max-num-seqs=1, SMEM
     pressure is determined by BLOCK, not batch token count)
   - gpu-memory-utilization: 0.9 → 0.95 (model uses ~17.5GB/GPU,
     KV cache for 100K tokens ≈ 1.38GB, plenty of headroom)

3. Added Dockerfile with patch_triton_tuning.py step.

4. Analysis document in optimizations/prefix_prefill_patch.py
   with full SMEM/register/occupancy derivation.
2026-07-30 15:33:44 +00:00

92 lines
3.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
patch_triton_tuning.py — BI-V100 Triton kernel parameter optimization
======================================================================
Patches prefix_prefill.py to use BI-V100-optimal BLOCK and NUM_WARPS values.
Hardware derivation:
BI-V100 SMEM = 48KB. Triton Flash Attention needs K+V tiles in SMEM:
SMEM = BLOCK_N × head_dim × sizeof(fp16) × 2
BLOCK_N=64, head_dim=128 → 32KB ≤ 48KB ✓ (current, correct)
BLOCK_N=128, head_dim=128 → 64KB > 48KB ✗ (would crash)
→ BLOCK must stay at 64.
NUM_WARPS derivation:
At BLOCK=64, each block does 64 query positions.
8 warps = 256 threads → each thread handles 32 elements from Q tile.
4 warps = 128 threads → each thread handles 64 elements.
With 50 SMs and typical grid of 37K+ blocks:
At 8 warps + 32KB SMEM: 1 block per SM (SMEM-limited)
At 4 warps + 32KB SMEM: potentially 2 blocks per SM
BI-V100 is bandwidth-limited (900 GB/s), not latency-limited.
Fewer warps hiding latency matters less; more blocks = better.
→ NUM_WARPS = 4
Deploy: python3 qwen3_6_scripts/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",
]
# Original line (baseline):
OLD_BLOCK = " BLOCK = 128 if current_platform.has_device_capability(80) else 64\n NUM_WARPS = 8"
# Optimized for BI-V100:
NEW_BLOCK = """\
# BI-V100 optimization (patch_triton_tuning.py):
# BLOCK=64: SMEM constraint — BLOCK_N=128 overflows 48KB at head_dim=128
# NUM_WARPS=4: bandwidth-limited GPU benefits from more blocks/SM over more warps
# Derivation: 4 warps at BLOCK=64 allows 2 concurrent blocks per SM,
# doubling occupancy vs 8 warps (which is SMEM-limited to 1 block/SM).
BLOCK = 64
NUM_WARPS = 4"""
def patch():
for path in PREFIX_PREFILL_PATHS:
if not os.path.exists(path):
continue
with open(path, "r") as f:
content = f.read()
if "NUM_WARPS = 4" in content:
print(f" [skip] {path}: already patched")
return
if OLD_BLOCK not in content:
# Try the alternative: maybe it's already using BLOCK=64 hardcoded
alt_old = " BLOCK = 64\n NUM_WARPS = 8"
if alt_old in content:
content = content.replace(alt_old, NEW_BLOCK, 1)
with open(path, "w") as f:
f.write(content)
print(f" [ok] {path}: patched NUM_WARPS 8→4 (BLOCK was already 64)")
return
print(f" [warn] {path}: original block not found, manual check needed")
return
content = content.replace(OLD_BLOCK, NEW_BLOCK, 1)
with open(path, "w") as f:
f.write(content)
print(f" [ok] {path}: patched BLOCK=64, NUM_WARPS=4")
return
print(" [error] prefix_prefill.py not found at any expected path")
def main():
print("=== patch_triton_tuning: BI-V100 Triton kernel optimization ===")
patch()
print("Done.")
if __name__ == "__main__":
main()