Download glint_metrics.py from SlayerLab/Slayer149: direct link, hf CLI and curl.
- Browser
- Download file 6.47 kB
-
https://huggingface.co/SlayerLab/Slayer149/resolve/main/glint_metrics.py
- Command line
-
hf download hf://SlayerLab/Slayer149/glint_metrics.py
-
curl -L -o glint_metrics.py https://huggingface.co/SlayerLab/Slayer149/resolve/main/glint_metrics.py
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} | |