Lottolabs's picture
Upload verified mixed BFP4/BFP8 checkpoint with MTP and evaluation evidence
12f320c verified
Raw History Blame Contribute Delete
6.01 kB
"""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()