"""Compare teacher-forced full-vocabulary distributions, never top-k approximations.""" import argparse import json import math from pathlib import Path import numpy as np def log_probs(logits): values = np.asarray(logits, dtype=np.float64) values -= values.max(axis=-1, keepdims=True) return values - np.log(np.exp(values).sum(axis=-1, keepdims=True)) def main(): parser = argparse.ArgumentParser() parser.add_argument('--reference', type=Path, required=True) parser.add_argument('--candidate', type=Path, required=True) parser.add_argument('--output', type=Path, required=True) parser.add_argument('--assistant-masks', type=Path, help='Verified assistant-target masks; omit for whole-sequence diagnostics') args = parser.parse_args() mask_info = json.loads(args.assistant_masks.read_text()) if args.assistant_masks else None for directory in (args.reference, args.candidate): metadata = json.loads((directory / 'metadata.json').read_text()) if metadata.get('status') != 'complete': raise ValueError(f'Incomplete inference artifacts: {directory}') if metadata.get('generation_max_new_tokens', 0): raise ValueError('Free-generation targets are not heldout teacher-forced data') reference = {r['id']: r for r in map(json.loads, (args.reference / 'records.jsonl').read_text().splitlines())} candidates = list(map(json.loads, (args.candidate / 'records.jsonl').read_text().splitlines())) if not candidates or len({r['id'] for r in candidates}) != len(candidates): raise ValueError('Missing or duplicate candidate records') results = [] token_kl_by_split = {} for candidate in candidates: ref = reference[candidate['id']] if ref['token_ids'] != candidate['token_ids'] or ref['split'] != candidate['split']: raise ValueError(f'Misaligned inputs: {candidate["id"]}') a = np.load(args.reference / ref['logits_file'], mmap_mode='r') b = np.load(args.candidate / candidate['logits_file'], mmap_mode='r') targets = np.asarray(ref['token_ids'][1:], dtype=np.int64) if a.shape != b.shape or a.ndim != 2 or a.shape[0] != len(targets): raise ValueError(f'Misaligned logits: {candidate["id"]}: {a.shape}, {b.shape}') target_mask = np.ones(len(targets), dtype=bool) domain = None if mask_info: mask_record = mask_info['records'][candidate['id']] if mask_record['token_ids'] != ref['token_ids'] or mask_record['split'] != ref['split']: raise ValueError('Assistant mask belongs to different input tokens or split') target_mask = np.asarray(mask_record['target_mask'], dtype=bool) domain = mask_record['domain'] if target_mask.shape != targets.shape or not target_mask.any(): raise ValueError('Missing or misaligned assistant targets') kl, ref_loss, quant_loss, agreements = [], [], [], [] for start in range(0, len(targets), 16): end = min(start + 16, len(targets)) keep = target_mask[start:end] if not keep.any(): continue lp, lq = log_probs(a[start:end][keep]), log_probs(b[start:end][keep]) if not np.isfinite(lp).all() or not np.isfinite(lq).all(): raise ValueError(f'Nonfinite logits: {candidate["id"]}') kl.extend(np.sum(np.exp(lp) * (lp - lq), axis=-1).tolist()) row = np.arange(int(keep.sum())) selected_targets = targets[start:end][keep] ref_loss.extend((-lp[row, selected_targets]).tolist()) quant_loss.extend((-lq[row, selected_targets]).tolist()) agreements.extend((lp.argmax(axis=-1) == lq.argmax(axis=-1)).tolist()) results.append({'id': candidate['id'], 'split': candidate['split'], 'domain': domain, 'tokens': len(kl), 'kl_sum': sum(kl), 'kl_mean': float(np.mean(kl)), 'kl_p95': float(np.percentile(kl, 95)), 'reference_nll_sum': sum(ref_loss), 'candidate_nll_sum': sum(quant_loss), 'top1_agreements': sum(agreements)}) token_kl_by_split.setdefault(candidate['split'], []).extend(kl) if domain: token_kl_by_split.setdefault(candidate['split'] + '/' + domain, []).extend(kl) results[-1].update({'kl_p99_9': float(np.percentile(kl, 99.9)), 'kl_max': max(kl)}) summaries = {} for split in sorted(token_kl_by_split): rows = [r for r in results if r['split'] == split or (r['domain'] and r['split'] + '/' + r['domain'] == split)] n = sum(r['tokens'] for r in rows) ref_nll = sum(r['reference_nll_sum'] for r in rows) / n cand_nll = sum(r['candidate_nll_sum'] for r in rows) / n summaries[split] = {'records': len(rows), 'tokens': n, 'kl_mean': sum(r['kl_sum'] for r in rows) / n, 'reference_nll': ref_nll, 'candidate_nll': cand_nll, 'nll_delta': cand_nll - ref_nll, 'reference_perplexity': math.exp(ref_nll), 'candidate_perplexity': math.exp(cand_nll), 'top1_agreement': sum(r['top1_agreements'] for r in rows) / n} values = token_kl_by_split[split] summaries[split].update({'kl_median': float(np.median(values)), 'kl_p95': float(np.percentile(values, 95)), 'kl_p99': float(np.percentile(values, 99)), 'kl_p99_9': float(np.percentile(values, 99.9)), 'kl_max': max(values)}) result = {'metric': 'KL(reference || candidate), natural logarithms, all vocabulary entries, teacher-forced positions', 'reference': str(args.reference), 'candidate': str(args.candidate), 'splits': summaries, 'records': results} result['position_scope'] = 'assistant outputs, tool calls and EOS; supplied headers excluded' if mask_info else 'all sequence targets including user/tool prompts' result['assistant_masks'] = str(args.assistant_masks) if args.assistant_masks else None args.output.write_text(json.dumps(result, indent=2) + '\n') print(json.dumps(summaries, indent=2)) if __name__ == '__main__': main()