from __future__ import annotations import argparse import time import torch import torch.nn.functional as F from tokenizers import Tokenizer from model import LogosModel def infer_config(sd: dict) -> dict: vocab_size, dim = sd['embed_tokens.weight'].shape n_layers = 1 + max((int(k.split('.')[1]) for k in sd if k.startswith('layers.'))) head_dim = sd['layers.0.attn.q_norm.weight'].shape[0] n_heads = dim // head_dim n_kv_heads = sd['layers.0.attn.k_proj.weight'].shape[0] // head_dim mlp_hidden = sd['layers.0.mlp.gate_proj.weight'].shape[0] return dict(vocab_size=vocab_size, dim=dim, n_layers=n_layers, n_heads=n_heads, n_kv_heads=n_kv_heads, mlp_hidden=mlp_hidden) def load_model(ckpt_path: str, device: str) -> LogosModel: state = torch.load(ckpt_path, map_location='cpu') sd = state['model'] if 'model' in state else state cfg = infer_config(sd) print(f'[eval] checkpoint architecture: {cfg}', flush=True) model = LogosModel(**cfg) model.load_state_dict(sd) model.eval() model.to(device) return model class _Encoding: def __init__(self, ids): self.ids = ids class _HFTokenizerAdapter: def __init__(self, hf_tok): self.hf_tok = hf_tok def encode(self, text): return _Encoding(self.hf_tok.encode(text, add_special_tokens=False)) class _HFModelAdapter: def __init__(self, hf_model): self.hf_model = hf_model def eval(self): self.hf_model.eval() return self def to(self, device): self.hf_model.to(device) return self def __call__(self, input_ids): return (self.hf_model(input_ids=input_ids).logits, None) def _join(context: str, choice: str) -> str: if context and (not context[-1].isspace()) and choice and (not choice[0].isspace()): return context + ' ' + choice return context + choice @torch.no_grad() def _score_batch(model, tok, device: str, pairs: list[tuple[str, str]]) -> list[tuple[float, int]]: metas = [] for context, choice in pairs: ctx_ids = tok.encode(context).ids whole_ids = tok.encode(_join(context, choice)).ids if len(whole_ids) <= len(ctx_ids): whole_ids = ctx_ids + tok.encode(choice).ids metas.append((whole_ids, len(whole_ids) - len(ctx_ids))) max_len = max((len(w) for w, _ in metas)) batch_ids = torch.zeros((len(metas), max_len), dtype=torch.long) for i, (whole_ids, _) in enumerate(metas): batch_ids[i, :len(whole_ids)] = torch.tensor(whole_ids, dtype=torch.long) batch_ids = batch_ids.to(device) logits, _ = model(batch_ids) logprobs = F.log_softmax(logits.float(), dim=-1) scores = [] for i, (whole_ids, cont_len) in enumerate(metas): start = len(whole_ids) - cont_len total = 0.0 for j in range(cont_len): pos = start - 1 + j tid = whole_ids[start + j] total += logprobs[i, pos, tid].item() scores.append((total, max(cont_len, 1))) return scores def batched_score(model, tok, device, pairs: list[tuple[str, str]], batch_size: int, tag: str='') -> list[tuple[float, int]]: scores: list[tuple[float, int]] = [(0.0, 1)] * len(pairs) n_batches = (len(pairs) + batch_size - 1) // batch_size t0 = time.time() for bi, start in enumerate(range(0, len(pairs), batch_size)): chunk = pairs[start:start + batch_size] chunk_scores = _score_batch(model, tok, device, chunk) scores[start:start + len(chunk)] = chunk_scores if tag and (bi % 10 == 0 or bi == n_batches - 1): elapsed = time.time() - t0 done = start + len(chunk) rate = done / elapsed if elapsed > 0 else 0 print(f'[eval] {tag}: {done}/{len(pairs)} pairs scored ({rate:.1f}/s, {elapsed:.0f}s elapsed)', flush=True) return scores def _pick(ex_scores, normalized: bool) -> int: if normalized: return max(range(len(ex_scores)), key=lambda k: ex_scores[k][0] / ex_scores[k][1]) return max(range(len(ex_scores)), key=lambda k: ex_scores[k][0]) def eval_hellaswag_style(model, tok, device, ds, ctx_key, endings_key, label_key, limit=None, batch_size=32, tag=''): if limit: ds = ds.select(range(min(limit, len(ds)))) pairs, counts = ([], []) for row in ds: endings = row[endings_key] for e in endings: pairs.append((row[ctx_key], e)) counts.append(len(endings)) scores = batched_score(model, tok, device, pairs, batch_size, tag) raw_correct, norm_correct, idx = (0, 0, 0) for i, row in enumerate(ds): n = counts[i] ex_scores = scores[idx:idx + n] idx += n gold = int(row[label_key]) raw_correct += int(_pick(ex_scores, False) == gold) norm_correct += int(_pick(ex_scores, True) == gold) return (norm_correct / len(ds), len(ds), raw_correct / len(ds)) def eval_arc(model, tok, device, config, limit=None, batch_size=32, tag=''): from datasets import load_dataset ds = load_dataset('allenai/ai2_arc', config, split='test') if limit: ds = ds.select(range(min(limit, len(ds)))) kept_rows = [] pairs, counts = ([], []) for row in ds: choices = row['choices']['text'] labels = row['choices']['label'] answer = row['answerKey'] if answer not in labels or not choices: continue kept_rows.append((row, labels.index(answer))) for c in choices: pairs.append((row['question'], c)) counts.append(len(choices)) scores = batched_score(model, tok, device, pairs, batch_size, tag) raw_correct, norm_correct, idx = (0, 0, 0) for (row, gold_idx), n in zip(kept_rows, counts): ex_scores = scores[idx:idx + n] idx += n raw_correct += int(_pick(ex_scores, False) == gold_idx) norm_correct += int(_pick(ex_scores, True) == gold_idx) return (norm_correct / len(kept_rows), len(kept_rows), raw_correct / len(kept_rows)) def eval_piqa(model, tok, device, limit=None, batch_size=32, tag=''): from datasets import load_dataset ds = load_dataset('ybisk/piqa', split='validation', revision='refs/convert/parquet') if limit: ds = ds.select(range(min(limit, len(ds)))) pairs = [] for row in ds: pairs.append((row['goal'], row['sol1'])) pairs.append((row['goal'], row['sol2'])) scores = batched_score(model, tok, device, pairs, batch_size, tag) raw_correct, norm_correct = (0, 0) for i, row in enumerate(ds): pair = [scores[2 * i], scores[2 * i + 1]] gold = int(row['label']) raw_correct += int(_pick(pair, False) == gold) norm_correct += int(_pick(pair, True) == gold) return (norm_correct / len(ds), len(ds), raw_correct / len(ds)) def N(score_pct: float, chance: float) -> float: return 100 * (score_pct - chance) / (100 - chance) def main(): parser = argparse.ArgumentParser() parser.add_argument('--ckpt', default=None) parser.add_argument('--hf_model', default=None, help='if set, evaluate a HF transformers reference model instead of a LogosModel checkpoint -- used to calibrate the harness against a model with known published scores') parser.add_argument('--device', default='cpu') parser.add_argument('--threads', type=int, default=24) parser.add_argument('--batch_size', type=int, default=32) parser.add_argument('--limit', type=int, default=None, help='cap per-benchmark examples (debug/speed)') args = parser.parse_args() if not args.ckpt and (not args.hf_model): parser.error('one of --ckpt or --hf_model is required') torch.set_num_threads(args.threads) if args.hf_model: from transformers import AutoModelForCausalLM, AutoTokenizer print(f'[eval] loading HF reference model {args.hf_model}...', flush=True) tok = _HFTokenizerAdapter(AutoTokenizer.from_pretrained(args.hf_model)) model = _HFModelAdapter(AutoModelForCausalLM.from_pretrained(args.hf_model)).eval().to(args.device) ckpt_label = args.hf_model else: tok = Tokenizer.from_file('artifacts/tokenizer.json') print(f'[eval] loading checkpoint {args.ckpt}...', flush=True) model = load_model(args.ckpt, args.device) ckpt_label = args.ckpt results = {} from datasets import load_dataset t0 = time.time() print('[eval] running HellaSwag...', flush=True) hs_ds = load_dataset('Rowan/hellaswag', split='validation') results['hellaswag'] = eval_hellaswag_style(model, tok, args.device, hs_ds, 'ctx', 'endings', 'label', args.limit, args.batch_size, tag='hellaswag') print(f"[eval] HellaSwag acc_norm: {results['hellaswag'][0]:.4f} (n={results['hellaswag'][1]}, {time.time() - t0:.0f}s)", flush=True) t0 = time.time() print('[eval] running ARC-Easy...', flush=True) results['arc_easy'] = eval_arc(model, tok, args.device, 'ARC-Easy', args.limit, args.batch_size, tag='arc_easy') print(f"[eval] ARC-Easy acc_norm: {results['arc_easy'][0]:.4f} (n={results['arc_easy'][1]}, {time.time() - t0:.0f}s)", flush=True) t0 = time.time() print('[eval] running ARC-Challenge...', flush=True) results['arc_challenge'] = eval_arc(model, tok, args.device, 'ARC-Challenge', args.limit, args.batch_size, tag='arc_challenge') print(f"[eval] ARC-Challenge acc_norm: {results['arc_challenge'][0]:.4f} (n={results['arc_challenge'][1]}, {time.time() - t0:.0f}s)", flush=True) t0 = time.time() print('[eval] running PIQA...', flush=True) results['piqa'] = eval_piqa(model, tok, args.device, args.limit, args.batch_size, tag='piqa') print(f"[eval] PIQA acc_norm: {results['piqa'][0]:.4f} (n={results['piqa'][1]}, {time.time() - t0:.0f}s)", flush=True) t0 = time.time() print('[eval] running ArithMark-3.0...', flush=True) am_ds = load_dataset('AxiomicLabs/ArithMark-3.0', split='train') results['arithmark3'] = eval_hellaswag_style(model, tok, args.device, am_ds, 'ctx', 'endings', 'label', args.limit, args.batch_size, tag='arithmark3') print(f"[eval] ArithMark-3.0 acc_norm: {results['arithmark3'][0]:.4f} (n={results['arithmark3'][1]}, {time.time() - t0:.0f}s)", flush=True) def index_from(hs, arc_e, arc_c, piqa, am): return (N(hs * 100, 25) + N((arc_e + arc_c) / 2 * 100, 25) + N(piqa * 100, 50) + 0.65 * N(am * 100, 25)) / 3.65 print('', flush=True) print('[eval] ==== acc vs acc_norm ====', flush=True) print(f"[eval] {'task':<16}{'acc':>9}{'acc_norm':>11}", flush=True) for key, label in [('hellaswag', 'HellaSwag'), ('arc_easy', 'ARC-Easy'), ('arc_challenge', 'ARC-Chall'), ('piqa', 'PIQA'), ('arithmark3', 'ArithMark-3')]: norm, _n, raw = results[key] print(f'[eval] {label:<16}{raw * 100:>8.2f}%{norm * 100:>10.2f}%', flush=True) idx_raw = index_from(results['hellaswag'][2], results['arc_easy'][2], results['arc_challenge'][2], results['piqa'][2], results['arithmark3'][2]) idx_norm = index_from(results['hellaswag'][0], results['arc_easy'][0], results['arc_challenge'][0], results['piqa'][0], results['arithmark3'][0]) print(f'[eval] Index from acc: {idx_raw:.2f}', flush=True) print(f'[eval] Index from acc_norm: {idx_norm:.2f}', flush=True) hs_acc, hs_n = (results['hellaswag'][0], results['hellaswag'][1]) e_acc, e_n = (results['arc_easy'][0], results['arc_easy'][1]) c_acc, c_n = (results['arc_challenge'][0], results['arc_challenge'][1]) piqa_acc, piqa_n = (results['piqa'][0], results['piqa'][1]) am_acc, am_n = (results['arithmark3'][0], results['arithmark3'][1]) combined_arc = (e_acc + c_acc) / 2 idx = (N(hs_acc * 100, 25) + N(combined_arc * 100, 25) + N(piqa_acc * 100, 50) + 0.65 * N(am_acc * 100, 25)) / 3.65 print('', flush=True) print('[eval] ==== SUMMARY ====', flush=True) print(f'[eval] checkpoint: {ckpt_label}', flush=True) print(f'[eval] HellaSwag: {hs_acc * 100:.2f}% (n={hs_n})', flush=True) print(f'[eval] ARC-Easy: {e_acc * 100:.2f}% (n={e_n})', flush=True) print(f'[eval] ARC-Challenge: {c_acc * 100:.2f}% (n={c_n})', flush=True) print(f'[eval] Combined ARC: {combined_arc * 100:.2f}%', flush=True) print(f'[eval] PIQA: {piqa_acc * 100:.2f}% (n={piqa_n})', flush=True) print(f'[eval] ArithMark-3.0: {am_acc * 100:.2f}% (n={am_n})', flush=True) print(f'[eval] Intelligence Index (approx): {idx:.2f}', flush=True) print(f'[eval] target (GPT-X2.5-135M): 25.17', flush=True) if __name__ == '__main__': main()