"""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()