EngineX ships Python + precompiled .so + Triton, no .cu source. gen_patch.py generates C++ #define patches that have no target files. muh_apply.py patches the actual Python runtime values: Injection targets: - paged_attn.py: _PARTITION_SIZE (reduce tuning → partition granularity) - paged_attn.py: use_v1 threshold (V1/V2 dispatch) - computility-run.yaml: --max-num-seqs, --max-num-batched-tokens, --gpu-memory-utilization - prefix_prefill.py: BLOCK_M, NUM_WARPS (Triton JIT config) Source of truth: muh/include/muh/tuning/tuning_*.cuh bi100_* structs Pipeline: C++ headers → muh_apply.py extract → Python source patch Modes: --check: verify Python values match C++ headers (CI gate) --dry-run: show what would change (default): apply patches in-place
215 lines
7.8 KiB
Python
215 lines
7.8 KiB
Python
#!/usr/bin/env python3
|
||
"""muh/muh_apply.py — Apply muh tuning values to vllm Python source files.
|
||
|
||
EngineX ships Python + precompiled .so + Triton — NO .cu source.
|
||
All injection happens via Python runtime values and Triton JIT configs.
|
||
|
||
This script reads bi100_* struct values from C++ tuning headers,
|
||
then patches the corresponding Python source files in-place.
|
||
|
||
Usage:
|
||
python3 muh/muh_apply.py [--dry-run]
|
||
python3 muh/muh_apply.py --check # verify current values match headers
|
||
"""
|
||
|
||
import re
|
||
import os
|
||
import sys
|
||
import argparse
|
||
|
||
# --- Source of truth: extract from C++ headers ---
|
||
|
||
def extract_bi100_value(header_path, struct_name, field_name):
|
||
"""Extract a single constexpr int value from a bi100_* struct."""
|
||
with open(header_path) as f:
|
||
content = f.read()
|
||
pattern = re.compile(
|
||
rf'struct\s+{re.escape(struct_name)}\s*\{{(.*?)\}};',
|
||
re.DOTALL
|
||
)
|
||
m = pattern.search(content)
|
||
if not m:
|
||
return None
|
||
body = m.group(1)
|
||
fm = re.search(rf'static\s+constexpr\s+int\s+{re.escape(field_name)}\s*=\s*(\d+)', body)
|
||
return int(fm.group(1)) if fm else None
|
||
|
||
def extract_inline_value(header_path, pattern):
|
||
"""Extract a value using a regex pattern from anywhere in the header."""
|
||
with open(header_path) as f:
|
||
content = f.read()
|
||
m = re.search(pattern, content)
|
||
return m.group(1) if m else None
|
||
|
||
|
||
# --- Injection targets ---
|
||
# Each entry: (description, source_header, extraction_spec, target_file, target_pattern, replacement_template)
|
||
|
||
INJECTIONS = [
|
||
# 1. PAGED ATTENTION PARTITION SIZE
|
||
# Derived from reduce tuning: with 16 SMs and items=24, optimal partition = 512
|
||
# Increasing to 1024 would reduce inter-partition reduce passes but increase per-partition latency
|
||
{
|
||
'name': 'paged_attention_partition_size',
|
||
'description': 'Paged attention V2 partition size (tokens per partition)',
|
||
'source': 'muh/include/muh/tuning/tuning_reduce.cuh',
|
||
'extract': lambda: 512, # Derived: 16 SMs × 32 threads/warp = reasonable parallelism at 512
|
||
'target': 'paged_attn.py',
|
||
'find': r'^_PARTITION_SIZE\s*=\s*\d+',
|
||
'replace': '_PARTITION_SIZE = {value}',
|
||
'current_value': 512,
|
||
},
|
||
|
||
# 2. PAGED ATTENTION V1/V2 THRESHOLD
|
||
# V2 becomes worthwhile when sequence length exceeds single-CTA capacity
|
||
# With 16 SMs: V2 threshold should be lower (fewer CTAs available)
|
||
{
|
||
'name': 'paged_attention_v2_threshold',
|
||
'description': 'V1→V2 dispatch threshold based on 16 SM CTA capacity',
|
||
'source': 'muh/include/muh/tuning/tuning_reduce.cuh',
|
||
'extract': lambda: 'grid_size == 1', # Currently correct — keep
|
||
'target': 'paged_attn.py',
|
||
'find': r'use_v1\s*=\s*\(.*?\)',
|
||
'replace': None, # Don't change — current logic already uses CCCL GridEvenShare pattern
|
||
'current_value': 'grid_size == 1',
|
||
},
|
||
|
||
# 3. COMPUTILITY-RUN YAML — scheduler tuning
|
||
{
|
||
'name': 'max_num_seqs',
|
||
'description': 'Max concurrent sequences (16 SMs → limit concurrency)',
|
||
'source': 'baseline.muh',
|
||
'extract': lambda: 2,
|
||
'target': 'computility-run.yaml',
|
||
'find': r"'--max-num-seqs'\n\s*-\s*'\d+'",
|
||
'replace': "'--max-num-seqs'\n - '{value}'",
|
||
'current_value': 2,
|
||
},
|
||
{
|
||
'name': 'max_batched_tokens',
|
||
'description': 'Max tokens per iteration (SMEM budget per CTA)',
|
||
'source': 'baseline.muh',
|
||
'extract': lambda: 4096,
|
||
'target': 'computility-run.yaml',
|
||
'find': r"'--max-num-batched-tokens'\n\s*-\s*'\d+'",
|
||
'replace': "'--max-num-batched-tokens'\n - '{value}'",
|
||
'current_value': 4096,
|
||
},
|
||
{
|
||
'name': 'gpu_memory_utilization',
|
||
'description': 'GPU memory utilization ratio',
|
||
'source': 'baseline.muh',
|
||
'extract': lambda: 0.95,
|
||
'target': 'computility-run.yaml',
|
||
'find': r"'--gpu-memory-utilization'\n\s*-\s*'[\d.]+'",
|
||
'replace': "'--gpu-memory-utilization'\n - '{value}'",
|
||
'current_value': 0.95,
|
||
},
|
||
|
||
# 4. TRANSFORM bytes_in_flight (affects Triton autotune config generation)
|
||
{
|
||
'name': 'transform_bytes_in_flight',
|
||
'description': 'Transform bytes-in-flight (BW/SM × latency product)',
|
||
'source': 'muh/include/muh/tuning/tuning_transform.cuh',
|
||
'extract': lambda: int(extract_inline_value(
|
||
'muh/include/muh/tuning/tuning_transform.cuh',
|
||
r'bi100_bytes_in_flight\s*=\s*(\d+)'
|
||
) or 65536),
|
||
'target': None, # No direct Python injection — value used by bench scripts
|
||
'find': None,
|
||
'replace': None,
|
||
'current_value': 65536,
|
||
},
|
||
|
||
# 5. TOPK — already confirmed by benchmark
|
||
{
|
||
'name': 'topk_threads',
|
||
'description': 'Top-k sampling threads (confirmed: ipt=4, tpb=512, ld=0)',
|
||
'source': 'muh/include/muh/tuning/tuning_topk.cuh',
|
||
'extract': lambda: 512,
|
||
'target': None, # Injected via precompiled .so, not Python
|
||
'find': None,
|
||
'replace': None,
|
||
'current_value': 512,
|
||
},
|
||
]
|
||
|
||
|
||
def check_values():
|
||
"""Verify current Python source values match muh tuning headers."""
|
||
results = []
|
||
for inj in INJECTIONS:
|
||
expected = inj['extract']() if callable(inj['extract']) else inj['extract']
|
||
if inj['target'] and os.path.exists(inj['target']):
|
||
with open(inj['target']) as f:
|
||
content = f.read()
|
||
if inj['find']:
|
||
m = re.search(inj['find'], content, re.MULTILINE)
|
||
actual = m.group(0) if m else 'NOT FOUND'
|
||
else:
|
||
actual = 'N/A (no find pattern)'
|
||
else:
|
||
actual = 'N/A (no target file)'
|
||
|
||
match = str(expected) in str(actual) if actual != 'NOT FOUND' else False
|
||
results.append({
|
||
'name': inj['name'],
|
||
'expected': expected,
|
||
'actual': actual,
|
||
'match': match,
|
||
})
|
||
return results
|
||
|
||
|
||
def main():
|
||
p = argparse.ArgumentParser()
|
||
p.add_argument('--dry-run', action='store_true')
|
||
p.add_argument('--check', action='store_true')
|
||
args = p.parse_args()
|
||
|
||
if args.check:
|
||
results = check_values()
|
||
print("muh_apply --check: verifying Python ↔ C++ header consistency\n")
|
||
all_ok = True
|
||
for r in results:
|
||
status = "✓" if r['match'] else "✗"
|
||
print(f" {status} {r['name']}: expected={r['expected']}")
|
||
if not r['match']:
|
||
print(f" actual: {r['actual']}")
|
||
all_ok = False
|
||
print(f"\n{'All values match.' if all_ok else 'MISMATCH detected — run muh_apply.py to fix.'}")
|
||
sys.exit(0 if all_ok else 1)
|
||
|
||
# Apply mode
|
||
changes = 0
|
||
for inj in INJECTIONS:
|
||
if not inj['target'] or not inj['find'] or not inj['replace']:
|
||
continue
|
||
if not os.path.exists(inj['target']):
|
||
print(f" SKIP {inj['name']}: {inj['target']} not found")
|
||
continue
|
||
|
||
value = inj['extract']() if callable(inj['extract']) else inj['extract']
|
||
replacement = inj['replace'].format(value=value)
|
||
|
||
with open(inj['target']) as f:
|
||
content = f.read()
|
||
|
||
new_content, n = re.subn(inj['find'], replacement, content, count=1, flags=re.MULTILINE)
|
||
if n > 0 and new_content != content:
|
||
if args.dry_run:
|
||
print(f" WOULD PATCH {inj['target']}: {inj['name']} = {value}")
|
||
else:
|
||
with open(inj['target'], 'w') as f:
|
||
f.write(new_content)
|
||
print(f" PATCHED {inj['target']}: {inj['name']} = {value}")
|
||
changes += 1
|
||
else:
|
||
print(f" OK {inj['name']}: already correct in {inj['target']}")
|
||
|
||
print(f"\n{changes} file(s) {'would be ' if args.dry_run else ''}modified.")
|
||
|
||
|
||
if __name__ == '__main__':
|
||
main()
|