glint_parity_eval: arch-aware loader (Qwen3 cfg from ckpt) - fix 64M maintainer load
Browse files- glint_parity_eval.py +28 -20
glint_parity_eval.py
CHANGED
|
@@ -180,12 +180,14 @@ def evaluate_arc_easy(logits_fn, tokenizer, device):
|
|
| 180 |
# WIRING NASZEGO MODELU (Monter: wypelnij) -- to jedyna czesc nie-Glint.
|
| 181 |
# ---------------------------------------------------------------------------
|
| 182 |
def load_our_model(ckpt_path, tokenizer_path, device):
|
| 183 |
-
"""Zwroc (logits_fn, tokenizer).
|
| 184 |
-
|
| 185 |
-
|
|
|
|
|
|
|
| 186 |
import importlib.util, os
|
| 187 |
-
|
| 188 |
-
|
| 189 |
gpt_src = None
|
| 190 |
for cand in ("/workspace/gollem/corpus/scripts/train_gpt_ref.py",
|
| 191 |
os.path.join(os.path.dirname(os.path.abspath(__file__)), "train_gpt_ref.py"),
|
|
@@ -199,23 +201,29 @@ def load_our_model(ckpt_path, tokenizer_path, device):
|
|
| 199 |
GPT = tgr.GPT
|
| 200 |
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
|
| 201 |
sd = ck["model"] if isinstance(ck, dict) and "model" in ck else ck
|
| 202 |
-
sd = {k.replace("_orig_mod.", ""): v for k, v in sd.items()}
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 211 |
block = sd["pos.weight"].shape[0]
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
model = GPT(int(vocab), int(n_layer), int(n_embd), int(h), int(block))
|
| 215 |
-
miss, unexp = model.load_state_dict(sd, strict=False)
|
| 216 |
-
assert not unexp and miss in ([], ["head.weight"]), (miss, unexp)
|
| 217 |
model.eval().to(device)
|
| 218 |
-
print(f"[load_our_model]
|
|
|
|
| 219 |
def logits_fn(ids):
|
| 220 |
out = model(ids)
|
| 221 |
logits = out[0] if isinstance(out, (tuple, list)) else out
|
|
|
|
| 180 |
# WIRING NASZEGO MODELU (Monter: wypelnij) -- to jedyna czesc nie-Glint.
|
| 181 |
# ---------------------------------------------------------------------------
|
| 182 |
def load_our_model(ckpt_path, tokenizer_path, device):
|
| 183 |
+
"""Zwroc (logits_fn, tokenizer). logits_fn(ids[B,T]) -> logits[B,T,REAL_VOCAB].
|
| 184 |
+
TODO Monter: zaimportuj nasza GPT-klase (z train-kodu gollem), zaladuj ckpt,
|
| 185 |
+
ustaw eval()+to(device). Nasz block=1024 > 256 chunki Glinta wiec forward OK.
|
| 186 |
+
Wazne: przytnij logits do realnego vocab (bez padded-vocab) jesli mamy padding.
|
| 187 |
+
Ponizej szkielet - dopasuj do naszej sygnatury forward()."""
|
| 188 |
import importlib.util, os
|
| 189 |
+
tokenizer = HFTokenizer.from_file(tokenizer_path) # BPE-12k tokenizer.json
|
| 190 |
+
# import naszej klasy GPT z train_gpt_ref.py (typowe lokalizacje: pod / lokalnie)
|
| 191 |
gpt_src = None
|
| 192 |
for cand in ("/workspace/gollem/corpus/scripts/train_gpt_ref.py",
|
| 193 |
os.path.join(os.path.dirname(os.path.abspath(__file__)), "train_gpt_ref.py"),
|
|
|
|
| 201 |
GPT = tgr.GPT
|
| 202 |
ck = torch.load(ckpt_path, map_location="cpu", weights_only=False)
|
| 203 |
sd = ck["model"] if isinstance(ck, dict) and "model" in ck else ck
|
| 204 |
+
sd = {k.replace("_orig_mod.", ""): v for k, v in sd.items()} # strip torch.compile
|
| 205 |
+
vocab, n_embd = sd["tok.weight"].shape
|
| 206 |
+
n_layer = 1 + max(int(k.split(".")[1]) for k in sd if k.startswith("blocks."))
|
| 207 |
+
import types
|
| 208 |
+
cfgd = ck.get("config") if isinstance(ck, dict) else None
|
| 209 |
+
if cfgd: # RB2 arch-aware ckpt: odtworz arch z zapisanego configu
|
| 210 |
+
cfg = types.SimpleNamespace(
|
| 211 |
+
norm=cfgd.get("norm", "layernorm"), norm_eps=cfgd.get("norm_eps", 1e-6),
|
| 212 |
+
pos=cfgd.get("pos", "learned"), rope_theta=cfgd.get("rope_theta", 10000.0),
|
| 213 |
+
ffn=cfgd.get("ffn", "gelu"), ffn_mult=cfgd.get("ffn_mult", 2.667),
|
| 214 |
+
value_residual=cfgd.get("value_residual", False), qk_norm=cfgd.get("qk_norm", False))
|
| 215 |
+
n_head = int(cfgd.get("n_head") or os.environ.get("N_HEAD", "6"))
|
| 216 |
+
block = int(cfgd.get("block") or (sd["pos.weight"].shape[0] if "pos.weight" in sd else 1024))
|
| 217 |
+
else: # legacy pre-RB2 ckpt: LayerNorm / learned-pos / GELU
|
| 218 |
+
cfg = types.SimpleNamespace(norm="layernorm", norm_eps=1e-6, pos="learned", rope_theta=10000.0,
|
| 219 |
+
ffn="gelu", ffn_mult=2.667, value_residual=False, qk_norm=False)
|
| 220 |
+
n_head = int(os.environ.get("N_HEAD", "6"))
|
| 221 |
block = sd["pos.weight"].shape[0]
|
| 222 |
+
model = GPT(int(vocab), int(n_layer), int(n_embd), int(n_head), int(block), cfg)
|
| 223 |
+
model.load_state_dict(sd, strict=True)
|
|
|
|
|
|
|
|
|
|
| 224 |
model.eval().to(device)
|
| 225 |
+
print(f"[load_our_model] vocab={vocab} L={n_layer} d={n_embd} h={n_head} block={block} "
|
| 226 |
+
f"norm={cfg.norm} pos={cfg.pos} ffn={cfg.ffn} vr={cfg.value_residual} dev={device}", flush=True)
|
| 227 |
def logits_fn(ids):
|
| 228 |
out = model(ids)
|
| 229 |
logits = out[0] if isinstance(out, (tuple, list)) else out
|