187 lines
6.6 KiB
Python
187 lines
6.6 KiB
Python
|
|
#!/usr/bin/env python3
|
|||
|
|
"""muh/bi100_triton_configs.py — Generate BI-V100 Triton autotune configs
|
|||
|
|
|
|||
|
|
Produces triton.Config entries optimized for BI-V100 (SM=16, SMEM≤48KB).
|
|||
|
|
These get inserted into triton_flash_attention.py's @triton.autotune decorator
|
|||
|
|
and prefix_prefill.py's BLOCK/NUM_WARPS selection.
|
|||
|
|
|
|||
|
|
Strategy (from CCCL tuning patterns):
|
|||
|
|
- SM=16 → fewer CTAs, each should do more work → prefer larger BLOCK_M
|
|||
|
|
- SMEM 48KB → limits BLOCK_M × head_dim × element_size
|
|||
|
|
- 900GB/s HBM but only 16 SMs → 56GB/s per SM → memory bound
|
|||
|
|
- waves_per_eu maps to CTA occupancy per SM
|
|||
|
|
|
|||
|
|
CCCL reference patterns (from tuning_reduce.cuh, tuning_scan.cuh):
|
|||
|
|
- reduce: tpb=512, ipt=24 → large tile, fewer CTAs
|
|||
|
|
- scan: tpb=384, ipt=22 → balanced between tile size and occupancy
|
|||
|
|
"""
|
|||
|
|
|
|||
|
|
import json
|
|||
|
|
|
|||
|
|
# BI-V100 constraints
|
|||
|
|
SM_COUNT = 16
|
|||
|
|
SMEM_LIMIT = 49152 # 48KB, TBD if actually 32KB
|
|||
|
|
WARP_SIZE = 32
|
|||
|
|
HEAD_DIM = 128 # Qwen3.6
|
|||
|
|
ELEM_SIZE_FP16 = 2
|
|||
|
|
ELEM_SIZE_FP32 = 4
|
|||
|
|
|
|||
|
|
def smem_estimate(block_m, block_n, head_dim=HEAD_DIM, elem=ELEM_SIZE_FP16):
|
|||
|
|
"""Estimate SMEM for flash attention tile."""
|
|||
|
|
q_tile = block_m * head_dim * elem
|
|||
|
|
k_tile = head_dim * block_n * elem
|
|||
|
|
v_tile = block_n * head_dim * elem
|
|||
|
|
acc_tile = block_m * head_dim * ELEM_SIZE_FP32 # fp32 accumulator
|
|||
|
|
return q_tile + k_tile + v_tile + acc_tile
|
|||
|
|
|
|||
|
|
def generate_flash_attn_configs():
|
|||
|
|
"""Generate triton.Config entries for triton_flash_attention.py."""
|
|||
|
|
configs = []
|
|||
|
|
|
|||
|
|
# Sweep: BLOCK_M × BLOCK_N × num_warps
|
|||
|
|
candidates = [
|
|||
|
|
# (BLOCK_M, BLOCK_N, num_warps, num_stages, PRE_LOAD_V, rationale)
|
|||
|
|
(64, 32, 4, 1, False, "SM=16 conservative: small tile, moderate parallelism"),
|
|||
|
|
(64, 64, 4, 1, False, "SM=16 balanced: symmetric tile"),
|
|||
|
|
(64, 64, 4, 1, True, "SM=16 balanced + V preload"),
|
|||
|
|
(128, 32, 4, 1, False, "SM=16 asymmetric: tall Q tile for decode-heavy"),
|
|||
|
|
(128, 64, 4, 1, False, "SM=16 medium: good Q coverage"),
|
|||
|
|
(128, 64, 8, 1, False, "SM=16 medium + more warps"),
|
|||
|
|
(32, 32, 2, 1, False, "SM=16 minimal: highest occupancy"),
|
|||
|
|
(32, 64, 4, 1, False, "SM=16 wide-K: good for long context"),
|
|||
|
|
(64, 128, 4, 1, False, "SM=16 wide-KV: maximizes KV reuse"),
|
|||
|
|
(256, 64, 8, 1, False, "SM=16 tall: few CTAs, large Q coverage"),
|
|||
|
|
]
|
|||
|
|
|
|||
|
|
for bm, bn, warps, stages, preload, rationale in candidates:
|
|||
|
|
smem = smem_estimate(bm, bn)
|
|||
|
|
if smem > SMEM_LIMIT:
|
|||
|
|
continue
|
|||
|
|
|
|||
|
|
# Check thread count is valid
|
|||
|
|
threads = warps * WARP_SIZE
|
|||
|
|
if threads > 1024:
|
|||
|
|
continue
|
|||
|
|
|
|||
|
|
configs.append({
|
|||
|
|
"BLOCK_M": bm,
|
|||
|
|
"BLOCK_N": bn,
|
|||
|
|
"waves_per_eu": max(1, SM_COUNT * 2 // max(1, (bm * bn) // 1024)),
|
|||
|
|
"PRE_LOAD_V": preload,
|
|||
|
|
"num_warps": warps,
|
|||
|
|
"num_stages": stages,
|
|||
|
|
"smem_est": smem,
|
|||
|
|
"smem_pct": round(smem / SMEM_LIMIT * 100),
|
|||
|
|
"rationale": rationale,
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
return configs
|
|||
|
|
|
|||
|
|
def generate_prefill_configs():
|
|||
|
|
"""Generate BLOCK/NUM_WARPS configs for prefix_prefill.py."""
|
|||
|
|
configs = []
|
|||
|
|
|
|||
|
|
for block_m in [16, 32, 64, 128]:
|
|||
|
|
for num_warps in [2, 4, 8]:
|
|||
|
|
# prefix_prefill uses BLOCK_M = BLOCK_N (symmetric)
|
|||
|
|
smem = smem_estimate(block_m, block_m)
|
|||
|
|
if smem > SMEM_LIMIT:
|
|||
|
|
continue
|
|||
|
|
|
|||
|
|
threads = num_warps * WARP_SIZE
|
|||
|
|
if threads > 1024:
|
|||
|
|
continue
|
|||
|
|
|
|||
|
|
configs.append({
|
|||
|
|
"BLOCK": block_m,
|
|||
|
|
"NUM_WARPS": num_warps,
|
|||
|
|
"smem_est": smem,
|
|||
|
|
"smem_pct": round(smem / SMEM_LIMIT * 100),
|
|||
|
|
})
|
|||
|
|
|
|||
|
|
return configs
|
|||
|
|
|
|||
|
|
def generate_moe_configs():
|
|||
|
|
"""Generate BLOCK_SIZE_M configs for fused_moe → ixformer."""
|
|||
|
|
# ixformer only reads BLOCK_SIZE_M. The heuristic in fused_moe.py uses:
|
|||
|
|
# M <= 16 → BLOCK_SIZE_M = 16
|
|||
|
|
# M <= 32 → BLOCK_SIZE_M = 32
|
|||
|
|
# M <= 64 → BLOCK_SIZE_M = 64
|
|||
|
|
# else → BLOCK_SIZE_M = 128
|
|||
|
|
# We test all powers of 2 that are valid
|
|||
|
|
return [
|
|||
|
|
{"BLOCK_SIZE_M": m, "rationale": f"M-tile={m}, affects expert batch granularity"}
|
|||
|
|
for m in [16, 32, 64, 128, 256]
|
|||
|
|
]
|
|||
|
|
|
|||
|
|
def generate_triton_config_code(configs):
|
|||
|
|
"""Generate Python code for triton.Config entries."""
|
|||
|
|
lines = []
|
|||
|
|
lines.append("# BI-V100 optimized configs (SM=16, SMEM≤48KB, 900GB/s)")
|
|||
|
|
lines.append("# Generated by muh/bi100_triton_configs.py")
|
|||
|
|
lines.append("# Insert into triton_flash_attention.py @triton.autotune configs=[]")
|
|||
|
|
lines.append("")
|
|||
|
|
|
|||
|
|
for c in configs:
|
|||
|
|
lines.append(f"triton.Config(")
|
|||
|
|
lines.append(f" {{")
|
|||
|
|
lines.append(f" \"BLOCK_M\": {c['BLOCK_M']},")
|
|||
|
|
lines.append(f" \"BLOCK_N\": {c['BLOCK_N']},")
|
|||
|
|
lines.append(f" \"waves_per_eu\": {c['waves_per_eu']},")
|
|||
|
|
lines.append(f" \"PRE_LOAD_V\": {c['PRE_LOAD_V']},")
|
|||
|
|
lines.append(f" }},")
|
|||
|
|
lines.append(f" num_stages={c['num_stages']},")
|
|||
|
|
lines.append(f" num_warps={c['num_warps']},")
|
|||
|
|
lines.append(f"), # SMEM≈{c['smem_est']}B ({c['smem_pct']}%) — {c['rationale']}")
|
|||
|
|
lines.append("")
|
|||
|
|
|
|||
|
|
return "\n".join(lines)
|
|||
|
|
|
|||
|
|
|
|||
|
|
if __name__ == "__main__":
|
|||
|
|
print("=" * 70)
|
|||
|
|
print("BI-V100 Triton Config Generator")
|
|||
|
|
print("=" * 70)
|
|||
|
|
|
|||
|
|
print("\n--- Flash Attention Configs ---")
|
|||
|
|
fa_configs = generate_flash_attn_configs()
|
|||
|
|
for c in fa_configs:
|
|||
|
|
print(f" BLOCK_M={c['BLOCK_M']:>3d} BLOCK_N={c['BLOCK_N']:>3d} "
|
|||
|
|
f"warps={c['num_warps']} stages={c['num_stages']} "
|
|||
|
|
f"preload={c['PRE_LOAD_V']!s:5s} "
|
|||
|
|
f"SMEM={c['smem_est']:>5d} ({c['smem_pct']:>2d}%)")
|
|||
|
|
|
|||
|
|
print(f"\n Total valid: {len(fa_configs)} configs")
|
|||
|
|
print(f"\n Code output:")
|
|||
|
|
print(generate_triton_config_code(fa_configs))
|
|||
|
|
|
|||
|
|
print("\n--- Prefill Configs ---")
|
|||
|
|
pf_configs = generate_prefill_configs()
|
|||
|
|
for c in pf_configs:
|
|||
|
|
print(f" BLOCK={c['BLOCK']:>3d} NUM_WARPS={c['NUM_WARPS']} "
|
|||
|
|
f"SMEM={c['smem_est']:>5d} ({c['smem_pct']:>2d}%)")
|
|||
|
|
|
|||
|
|
print(f"\n Total valid: {len(pf_configs)} configs")
|
|||
|
|
|
|||
|
|
print("\n--- MoE Configs ---")
|
|||
|
|
moe_configs = generate_moe_configs()
|
|||
|
|
for c in moe_configs:
|
|||
|
|
print(f" BLOCK_SIZE_M={c['BLOCK_SIZE_M']:>3d} — {c['rationale']}")
|
|||
|
|
|
|||
|
|
# Save all configs
|
|||
|
|
all_configs = {
|
|||
|
|
"flash_attn": fa_configs,
|
|||
|
|
"prefill": pf_configs,
|
|||
|
|
"moe": moe_configs,
|
|||
|
|
"hardware": {
|
|||
|
|
"sm_count": SM_COUNT,
|
|||
|
|
"smem_limit": SMEM_LIMIT,
|
|||
|
|
"head_dim": HEAD_DIM,
|
|||
|
|
},
|
|||
|
|
}
|
|||
|
|
|
|||
|
|
out_path = "muh/bi100_configs.json"
|
|||
|
|
with open(out_path, "w") as f:
|
|||
|
|
json.dump(all_configs, f, indent=2)
|
|||
|
|
print(f"\nSaved: {out_path}")
|