File size: 14,014 Bytes
75dffb0 45e26fb 75dffb0 45e26fb 75dffb0 437a1d4 75dffb0 437a1d4 75dffb0 437a1d4 7a7066f 437a1d4 75dffb0 437a1d4 75dffb0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 | #!/usr/bin/env python3
"""
glint_parity_eval.py - EXACT port of Glint-1.3/benchmark.py eval-protocol,
model-agnostic. Measures OUR checkpoint on the BOARD's protocol so recon-position
is defensible (patrz labvault .../90-Ewaluacja/EvalHarnessParity.md).
Protocol fidelity (verbatim z Glint-1.3/benchmark.py):
- BLiMP: 67 configs, split='train', clip-first-256-tokens, raw-sum-logprobs
(NO BOS, NO length-norm), acc = good_ll > bad_ll.
- ARC-Easy: ai2_arc/ARC-Easy/test, zero-shot, candidate = question+" "+choice,
score = LL(q+choice) - LL(q), RAW acc (nie acc_norm).
- WikiText-2: wikitext-2-raw-v1/test, " ".join(rows).strip(),
non-overlapping 256-token chunks, context-RESET per chunk,
ppl = exp(total_NLL / n_token_predictions) <-- TOKEN-PPL (nasz tokenizer),
NIE BPB. To jest board-input dla WikiScore.
WIRING (Monter): wypelnij load_our_model() ponizej - import naszej GPT-klasy,
zaladuj ckpt, zwroc (model, logits_fn, tokenizer). logits_fn(input_ids_LongTensor[B,T])
MUSI zwrocic logits[B,T,vocab] (tylko realne vocab, bez padded-vocab).
Reszta = protokol Glint bez zmian. Odpal: python glint_parity_eval.py <ckpt> <tokenizer.json>
"""
import math, json, sys, time
import torch
import torch.nn.functional as F
import numpy as np
from datasets import load_dataset, concatenate_datasets
from tokenizers import Tokenizer as HFTokenizer
# ---------------------------------------------------------------------------
# GLINT EVAL-LOGIC (verbatim, model-agnostic: uzywa logits_fn + tokenizer)
# ---------------------------------------------------------------------------
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 _tok_path():
import os
for p in ("/workspace/.cache/huggingface/token", os.path.expanduser("~/.cache/huggingface/token"),
"/mnt/c/Users/Maggio03/.cache/huggingface/token"):
if os.path.exists(p):
return open(p).read().strip()
return None
def _rows(repo, config, split):
"""Robust loader: pyarrow-parquet via hf_hub_download (omija datasets-5.x load_dataset URI-bug)."""
import os, pyarrow.parquet as pq
from huggingface_hub import hf_hub_download, list_repo_files
tk = _tok_path()
files = list_repo_files(repo, repo_type="dataset", token=tk)
def match(f):
if not f.endswith(".parquet"): return False
base = os.path.basename(f).lower()
if split not in base and ("/"+split+"/") not in ("/"+f.lower()): return False
# dokladny katalog <config>/ (podciag lapal np. "transitive" w "intransitive" -> 6 fenomenow BLiMP 2x)
if config is not None and f.split("/")[0] != config: return False
return True
cands = [f for f in files if match(f)]
rows = []
for f in sorted(cands):
p = hf_hub_download(repo, f, repo_type="dataset", token=tk)
rows.extend(pq.read_table(p).to_pylist())
if not rows:
raise RuntimeError(f"_rows: brak parquet dla {repo} config={config} split={split}; kandydaci={cands[:5]}")
return rows
def evaluate_wikitext2(logits_fn, tokenizer, device):
rows = _rows("Salesforce/wikitext", "wikitext-2-raw-v1", "test")
text = " ".join(r["text"] for r in rows).strip()
ppl = compute_perplexity(logits_fn, tokenizer, text, device)
return {"wikitext2_ppl": round(ppl, 4)}
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)"
cap = os.environ.get("BLIMP_SAMPLE")
if cap: # opcjonalna próbka dla szybkości CPU (zaznaczyć w notatce)
import random; random.seed(1337); random.shuffle(ds); ds = ds[:int(cap)]
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}
# ---------------------------------------------------------------------------
# WIRING NASZEGO MODELU (Monter: wypelnij) -- to jedyna czesc nie-Glint.
# ---------------------------------------------------------------------------
def load_our_model(ckpt_path, tokenizer_path, device):
"""Zwroc (logits_fn, tokenizer). logits_fn(ids[B,T]) -> logits[B,T,REAL_VOCAB].
TODO Monter: zaimportuj nasza GPT-klase (z train-kodu gollem), zaladuj ckpt,
ustaw eval()+to(device). Nasz block=1024 > 256 chunki Glinta wiec forward OK.
Wazne: przytnij logits do realnego vocab (bez padded-vocab) jesli mamy padding.
Ponizej szkielet - dopasuj do naszej sygnatury forward()."""
import importlib.util, os
tokenizer = HFTokenizer.from_file(tokenizer_path) # BPE-12k tokenizer.json
# import naszej klasy GPT z train_gpt_ref.py (typowe lokalizacje: pod / lokalnie)
gpt_src = None
for cand in ("/workspace/gollem/corpus/scripts/train_gpt_ref.py",
os.path.join(os.path.dirname(os.path.abspath(__file__)), "train_gpt_ref.py"),
"/mnt/c/Projekty/Slayer/train-bdh-25m/train_gpt_ref.py"):
if os.path.exists(cand):
gpt_src = cand; break
if gpt_src is None:
raise FileNotFoundError("train_gpt_ref.py (klasa GPT) nie znaleziony")
spec = importlib.util.spec_from_file_location("tgr_glint", gpt_src)
tgr = importlib.util.module_from_spec(spec); spec.loader.exec_module(tgr)
GPT = tgr.GPT
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
sd = ck["model"] if isinstance(ck, dict) and "model" in ck else ck
sd = {k.replace("_orig_mod.", ""): v for k, v in sd.items()} # strip torch.compile
vocab, n_embd = sd["tok.weight"].shape
n_layer = 1 + max(int(k.split(".")[1]) for k in sd if k.startswith("blocks."))
import types
cfgd = ck.get("config") if isinstance(ck, dict) else None
if cfgd: # RB2 arch-aware ckpt: odtworz arch z zapisanego configu
cfg = types.SimpleNamespace(
norm=cfgd.get("norm", "layernorm"), norm_eps=cfgd.get("norm_eps", 1e-6),
pos=cfgd.get("pos", "learned"), rope_theta=cfgd.get("rope_theta", 10000.0),
ffn=cfgd.get("ffn", "gelu"), ffn_mult=cfgd.get("ffn_mult", 2.667),
value_residual=cfgd.get("value_residual", False), qk_norm=cfgd.get("qk_norm", False))
n_head = int(cfgd.get("n_head") or os.environ.get("N_HEAD", "6"))
block = int(cfgd.get("block") or (sd["pos.weight"].shape[0] if "pos.weight" in sd else 1024))
else: # legacy pre-RB2 ckpt: LayerNorm / learned-pos / GELU
cfg = types.SimpleNamespace(norm="layernorm", norm_eps=1e-6, pos="learned", rope_theta=10000.0,
ffn="gelu", ffn_mult=2.667, value_residual=False, qk_norm=False)
n_head = int(os.environ.get("N_HEAD", "6"))
block = sd["pos.weight"].shape[0]
model = GPT(int(vocab), int(n_layer), int(n_embd), int(n_head), int(block), cfg)
model.load_state_dict(sd, strict=True)
model.eval().to(device)
print(f"[load_our_model] vocab={vocab} L={n_layer} d={n_embd} h={n_head} block={block} "
f"norm={cfg.norm} pos={cfg.pos} ffn={cfg.ffn} vr={cfg.value_residual} dev={device}", flush=True)
def logits_fn(ids):
out = model(ids)
logits = out[0] if isinstance(out, (tuple, list)) else out
return logits[..., :tokenizer.get_vocab_size()]
return logits_fn, tokenizer
def main():
ckpt = sys.argv[1] if len(sys.argv) > 1 else "run_bpe16m_10b_e/ckpt.pt"
tok = sys.argv[2] if len(sys.argv) > 2 else "tokenizer.json"
device = "cuda" if torch.cuda.is_available() else "cpu"
logits_fn, tokenizer = load_our_model(ckpt, tok, device)
results = {}
print("1/3 WikiText-2 (token-PPL)...", flush=True)
results.update(evaluate_wikitext2(logits_fn, tokenizer, device))
print("2/3 BLiMP...", flush=True)
results.update(evaluate_blimp(logits_fn, tokenizer, device))
print("3/3 ARC-Easy...", flush=True)
results.update(evaluate_arc_easy(logits_fn, tokenizer, device))
print("GLINT-PROTOCOL RESULTS:", json.dumps(results, indent=2))
with open("glint_parity_results.json", "w") as f:
json.dump(results, f, indent=2)
if __name__ == "__main__":
main()
|