Slayer149 / glint_metrics.py
kacperwikiel's picture
Release GoLLeM 149M after 20B continuation tokens with full GLINT evaluation
3f431df verified
Raw History Blame Contribute Delete
6.47 kB
"""GLINT-1.3 likelihood protocol adapted to a logits callable; see NOTICE."""
import math
import torch
import torch.nn.functional as F
import numpy as np
def _rows(*args):
raise RuntimeError("Use evaluate_glint.py to install the pinned dataset loader")
def tokenize_many(tokenizer, texts, max_length=256):
all_ids = []
for text in texts:
ids = tokenizer.encode(text).ids
ids = [i for i in ids if i < tokenizer.get_vocab_size()]
if len(ids) > max_length:
ids = ids[:max_length]
all_ids.append(ids)
return all_ids
def batch_log_probs(logits_fn, tokenizer, texts, device, max_length=256, batch_size=128):
all_ids = tokenize_many(tokenizer, texts, max_length)
results = [-float("inf")] * len(all_ids)
with torch.inference_mode():
for start in range(0, len(all_ids), batch_size):
end = min(start + batch_size, len(all_ids))
batch = all_ids[start:end]
batch_indices = [j for j in range(start, end) if len(batch[j-start]) >= 2]
batch_seqs = [batch[j-start] for j in range(start, end) if len(batch[j-start]) >= 2]
if not batch_seqs:
continue
max_len = max(len(s) for s in batch_seqs)
B = len(batch_seqs)
padded_np = np.zeros((B, max_len - 1), dtype=np.int64)
targets_np = np.zeros((B, max_len - 1), dtype=np.int64)
mask_np = np.zeros((B, max_len - 1), dtype=bool)
for j, ids in enumerate(batch_seqs):
padded_np[j, :len(ids)-1] = ids[:-1]
targets_np[j, :len(ids)-1] = ids[1:]
mask_np[j, :len(ids)-1] = True
padded = torch.from_numpy(padded_np).to(device)
targets = torch.from_numpy(targets_np).to(device)
mask = torch.from_numpy(mask_np).to(device)
logits = logits_fn(padded)
log_probs = F.log_softmax(logits, dim=-1)
log_probs_flat = log_probs.view(-1, logits.size(-1))
targets_flat = targets.view(-1)
gathered = log_probs_flat[torch.arange(targets_flat.size(0), device=device), targets_flat]
gathered = gathered.view(B, -1)
gathered[~mask] = 0.0
sums = gathered.sum(dim=-1).tolist()
for bi, val in zip(batch_indices, sums):
results[bi] = val
return results
def compute_perplexity(logits_fn, tokenizer, text, device, max_length=256):
ids = tokenizer.encode(text).ids
ids = [i for i in ids if i < tokenizer.get_vocab_size()]
if len(ids) < 2:
return float("inf")
nll = 0.0; n_tokens = 0
for i in range(0, len(ids) - 1, max_length):
chunk = ids[i:i + max_length + 1]
if len(chunk) < 2:
continue
inputs = torch.tensor([chunk[:-1]], device=device)
targets = torch.tensor([chunk[1:]], device=device)
with torch.no_grad():
logits = logits_fn(inputs)
loss = F.cross_entropy(logits.view(-1, logits.size(-1)), targets.view(-1), reduction="sum")
nll += loss.item(); n_tokens += targets.numel()
return math.exp(nll / n_tokens) if n_tokens > 0 else float("inf")
BLIMP_CONFIGS = [
"adjunct_island","anaphor_gender_agreement","anaphor_number_agreement","animate_subject_passive",
"animate_subject_trans","causative","complex_NP_island","coordinate_structure_constraint_complex_left_branch",
"coordinate_structure_constraint_object_extraction","determiner_noun_agreement_1","determiner_noun_agreement_2",
"determiner_noun_agreement_irregular_1","determiner_noun_agreement_irregular_2","determiner_noun_agreement_with_adj_2",
"determiner_noun_agreement_with_adj_irregular_1","determiner_noun_agreement_with_adj_irregular_2",
"determiner_noun_agreement_with_adjective_1","distractor_agreement_relational_noun",
"distractor_agreement_relative_clause","drop_argument","ellipsis_n_bar_1","ellipsis_n_bar_2",
"existential_there_object_raising","existential_there_quantifiers_1","existential_there_quantifiers_2",
"existential_there_subject_raising","expletive_it_object_raising","inchoative","intransitive",
"irregular_past_participle_adjectives","irregular_past_participle_verbs","irregular_plural_subject_verb_agreement_1",
"irregular_plural_subject_verb_agreement_2","left_branch_island_echo_question","left_branch_island_simple_question",
"matrix_question_npi_licensor_present","npi_present_1","npi_present_2","only_npi_licensor_present","only_npi_scope",
"passive_1","passive_2","principle_A_c_command","principle_A_case_1","principle_A_case_2","principle_A_domain_1",
"principle_A_domain_2","principle_A_domain_3","principle_A_reconstruction","regular_plural_subject_verb_agreement_1",
"regular_plural_subject_verb_agreement_2","sentential_negation_npi_licensor_present","sentential_negation_npi_scope",
"sentential_subject_island","superlative_quantifiers_1","superlative_quantifiers_2","tough_vs_raising_1",
"tough_vs_raising_2","transitive","wh_island","wh_questions_object_gap","wh_questions_subject_gap",
"wh_questions_subject_gap_long_distance","wh_vs_that_no_gap","wh_vs_that_no_gap_long_distance",
"wh_vs_that_with_gap","wh_vs_that_with_gap_long_distance",
]
def evaluate_blimp(logits_fn, tokenizer, device):
import os
ds = []
for c in BLIMP_CONFIGS:
ds.extend(_rows("nyu-mll/blimp", c, "train"))
assert len(ds) == 67000, f"BLiMP: {len(ds)} par, oczekiwano 67000 (67 fenomenow x 1000)"
good = batch_log_probs(logits_fn, tokenizer, [e["sentence_good"] for e in ds], device)
bad = batch_log_probs(logits_fn, tokenizer, [e["sentence_bad"] for e in ds], device)
correct = sum(1 for g, b in zip(good, bad) if g > b)
return {"blimp_acc": round(correct/len(ds)*100, 2), "blimp_n": len(ds)}
def evaluate_arc_easy(logits_fn, tokenizer, device):
ds = _rows("allenai/ai2_arc", "ARC-Easy", "test")
correct = 0; total = 0
for ex in ds:
q = ex["question"]; ch = ex["choices"]
full = [q + " " + t for t in ch["text"]]
lps = batch_log_probs(logits_fn, tokenizer, full, device, batch_size=4)
lpq = batch_log_probs(logits_fn, tokenizer, [q], device)[0]
best = max(range(len(lps)), key=lambda j: lps[j] - lpq)
if ch["label"][best] == ex["answerKey"]:
correct += 1
total += 1
return {"arc_easy_acc": round(correct/total*100, 2), "arc_n": total}