Qwen-2.5-1B-RLCD-Fast / benchmark.py
epsilon3's picture
Make release M4-only and lead with speed and memory
728caeb verified
Raw History Blame Contribute Delete
8.86 kB
"""Synchronized paired comparison on real Qwen weights and upstream presets."""
import argparse
import gc
import json
import platform
import random
import statistics
import sys
import time
from pathlib import Path
import torch
import transformers
from transformers import AutoModelForCausalLM, AutoTokenizer
from tree_decode import FieldPlan, decode_fields, prefill
ROOT = Path(__file__).resolve().parent
SOURCE = ROOT / 'upstream' if (ROOT / 'upstream').exists() else ROOT
sys.path.insert(0, str(SOURCE))
from core.schema import StructuredSchema
def sync():
if torch.backends.mps.is_available():
torch.mps.synchronize()
def timed(fn):
sync()
start = time.perf_counter()
result = fn()
sync()
return result, (time.perf_counter() - start) * 1000
def summary(xs):
ordered = sorted(xs)
return dict(median_ms=statistics.median(xs), min_ms=min(xs),
p90_ms=ordered[min(len(xs)-1, int(len(xs)*0.9))], samples_ms=xs)
def speedup_interval(batch, tree):
rng = random.Random(83)
ratios = []
for _ in range(2000):
indices = [rng.randrange(len(batch)) for _ in batch]
ratios.append(statistics.median([batch[i] for i in indices]) /
statistics.median([tree[i] for i in indices]))
ratios.sort()
return [ratios[50], ratios[1949]]
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--model', default='Qwen/Qwen2.5-1.5B-Instruct')
ap.add_argument('--presets', nargs='+', default=['fintech_fraud', 'support_triage', 'code_security'])
ap.add_argument('--fields', nargs='+', type=int, default=[4, 16, 28])
ap.add_argument('--steps', type=int, default=1)
ap.add_argument('--repeats', type=int, default=12)
ap.add_argument('--warmup', type=int, default=3)
ap.add_argument('--output', default='results.json')
ap.add_argument('--device', choices=['auto', 'mps', 'cpu'], default='auto')
ap.add_argument('--dtype', choices=['auto', 'float16', 'float32'], default='auto')
args = ap.parse_args()
if args.repeats < 2 or args.warmup < 1:
ap.error('Use at least two repeats and one warmup')
torch.set_num_threads(4)
random.seed(42)
device = args.device
if device == 'auto':
device = 'mps' if torch.backends.mps.is_available() else 'cpu'
dtype = ({'float16': torch.float16, 'float32': torch.float32}.get(args.dtype)
or (torch.float16 if device == 'mps' else torch.float32))
tokenizer = AutoTokenizer.from_pretrained(args.model)
model = AutoModelForCausalLM.from_pretrained(args.model, dtype=dtype, attn_implementation='sdpa').to(device).eval()
print(f'Loaded {args.model} on {device} ({dtype})', flush=True)
report = {'environment': {'model': args.model, 'revision': model.config._commit_hash,
'torch': torch.__version__, 'transformers': transformers.__version__,
'gpu': device if device == 'mps' else None,
'platform': platform.platform(), 'dtype': str(dtype), 'attention': 'sdpa'},
'arguments': vars(args), 'cases': []}
for preset_name in args.presets:
preset = json.loads((SOURCE / 'presets' / f'{preset_name}.json').read_text(encoding='utf-8'))
for count in args.fields:
if count > len(preset['schema']):
continue
schema = StructuredSchema(dict(list(preset['schema'].items())[:count]))
print(f'Checking {preset_name}, {count} fields...', flush=True)
meta = schema.compile_parallel_metadata(tokenizer)
suffixes = [list(map(int, row[:length])) for row, length in zip(meta['suffixes_batch'], meta['suffix_lengths'])]
prompt = (f'<|im_start|>system\nClassify JSON attributes:\n{schema.to_parallel_schema_str()}<|im_end|>\n'
f'<|im_start|>user\n{preset["context"]}<|im_end|>\n<|im_start|>assistant\n{{\n')
ids = tokenizer.encode(prompt, return_tensors='pt').to(device)
plan = FieldPlan(suffixes, ids.shape[1], device, dtype, tokenizer.pad_token_id or 0)
cache = prefill(model, ids)
ref = decode_fields(model, cache, plan, 'batch', args.steps)
actual = decode_fields(model, cache, plan, 'tree', args.steps, ref.argmax(-1))
delta = (ref.float() - actual.float()).abs()
# FP16 kernels have different reduction orders for these shapes.
torch.testing.assert_close(actual, ref, atol=0.15 if dtype == torch.float16 else 1e-4,
rtol=0.01 if dtype == torch.float16 else 1e-4)
candidate_mismatches = []
for i, candidate_ids in enumerate(meta['cands_per_field']):
batch_choice = ref[:, i, candidate_ids].argmax(-1)
tree_choice = actual[:, i, candidate_ids].argmax(-1)
if not torch.equal(batch_choice, tree_choice):
candidate_mismatches.append(meta['field_items'][i][0])
candidates_equal = not candidate_mismatches
check = {'max_abs_logit_error': delta.max().item(), 'mean_abs_logit_error': delta.mean().item(),
'full_vocab_argmax_agreement': (ref.argmax(-1) == actual.argmax(-1)).float().mean().item(),
'candidate_decisions_equal': candidates_equal,
'candidate_decision_mismatches': candidate_mismatches,
'candidate_collision_fields': [name for (name, _), collision in zip(meta['field_items'], meta['has_collisions']) if collision]}
print(f'Correctness passed: {check}', flush=True)
del ref, actual, delta
funcs = {mode: (lambda m=mode: decode_fields(model, cache, plan, m, args.steps)) for mode in ('batch', 'tree')}
for _ in range(args.warmup):
for fn in funcs.values():
fn()
samples = {mode: [] for mode in funcs}
peaks = {}
for _ in range(args.repeats):
order = list(funcs)
random.shuffle(order)
for mode in order:
value, ms = timed(funcs[mode])
samples[mode].append(ms)
del value
print(f'Decode medians: batch={statistics.median(samples["batch"]):.1f} ms, tree={statistics.median(samples["tree"]):.1f} ms', flush=True)
for mode, fn in funcs.items():
if device == 'mps':
torch.mps.empty_cache()
baseline = torch.mps.driver_allocated_memory()
value, _ = timed(fn)
high_water = torch.mps.driver_allocated_memory()
peaks[mode] = {'driver_high_water_mib': high_water/2**20,
'extra_driver_mib': (high_water-baseline)/2**20}
del value
# Complete model request includes a fresh common prefill for each mode.
end_to_end = {mode: [] for mode in funcs}
del cache
for _ in range(args.repeats):
order = list(funcs)
random.shuffle(order)
for mode in order:
def request():
prefix = prefill(model, ids)
return decode_fields(model, prefix, plan, mode, args.steps)
value, ms = timed(request)
end_to_end[mode].append(ms)
del value
case = dict(preset=preset_name, fields=count, prefix_tokens=ids.shape[1],
suffix_tokens=sum(map(len, suffixes)), batch_padded_tokens=count*plan.width,
correctness=check, decode={m:summary(v) for m,v in samples.items()},
end_to_end={m:summary(v) for m,v in end_to_end.items()}, memory=peaks)
case['decode_speedup'] = statistics.median(samples['batch'])/statistics.median(samples['tree'])
case['end_to_end_speedup'] = statistics.median(end_to_end['batch'])/statistics.median(end_to_end['tree'])
case['decode_speedup_ci95'] = speedup_interval(samples['batch'], samples['tree'])
case['end_to_end_speedup_ci95'] = speedup_interval(end_to_end['batch'], end_to_end['tree'])
report['cases'].append(case)
Path(args.output).write_text(json.dumps(report, indent=2), encoding='utf-8')
print(f'{preset_name} {count} fields: decode {case["decode_speedup"]:.3f}x, end-to-end {case["end_to_end_speedup"]:.3f}x; {check}', flush=True)
gc.collect()
print(f'Results saved to {args.output}', flush=True)
if __name__ == '__main__':
main()