Download score_logits.py from Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150: direct link, hf CLI and curl.
- Browser
- Download file 6.01 kB
-
https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/score_logits.py
- Command line
-
hf download hf://Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/score_logits.py
-
curl -L -o score_logits.py https://huggingface.co/Lottolabs/Qwen3.5-9B-TT-Mixed-BFP4-BFP8-P150/resolve/main/score_logits.py
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() | |