cagliostro-v1 / eval_bench.py
TobiasLogic's picture
cagliostro-v1
5cec998 verified
Raw
History Blame Contribute Delete
12.6 kB
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()