"""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}