#!/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}")