diff options
Diffstat (limited to 'src/eval.py')
| -rw-r--r-- | src/eval.py | 90 |
1 files changed, 62 insertions, 28 deletions
diff --git a/src/eval.py b/src/eval.py index e93f95d..1cb1191 100644 --- a/src/eval.py +++ b/src/eval.py @@ -26,31 +26,46 @@ FLAGGED_SIEVE_PREDS = {121, 143, 169, 187, 209} def probe_report(model, cfg: Config, lo: int = 101, hi: int = 200) -> dict: - primes = sieve_primes(hi + 200) + try: + device = next(model.parameters()).device + except StopIteration: + device = torch.device("cpu") + margin = max(200, int(hi * 0.1) + 50) + primes = sieve_primes(hi + margin) correct = 0 errors = [] easy_total = 0 easy_wrong = 0 is_prime_task = cfg.task_mode == "is_prime" - for n in range(lo, hi + 1): - x = torch.tensor(encode_int(n, cfg), dtype=torch.long).unsqueeze(0) - gen = greedy_decode(model, x, cfg)[0].tolist() - pred = decode_tokens(gen, cfg) - if is_prime_task: - target = 1 if is_prime_n(n) else 0 - ok = (pred == 1) == (target == 1) # any non-"1" output reads as "composite" - else: - target = next_prime(n, primes) - ok = pred == target - is_easy = (n % 2 == 0) or (n % 5 == 0) # trivial composites (skip-evens / skip-5s) - if is_easy: - easy_total += 1 - if ok: - correct += 1 - else: - errors.append({"n": n, "target": target, "pred": pred}) + eval_bs = getattr(cfg, "eval_batch_size", 512) + inputs = list(range(lo, hi + 1)) + + for bi in range(0, len(inputs), eval_bs): + chunk = inputs[bi: bi + eval_bs] + xs = [encode_int(n, cfg) for n in chunk] + in_max = max(len(xi) for xi in xs) + x_tensor = torch.full((len(chunk), in_max), cfg.pad_id, dtype=torch.long, device=device) + for i, xi in enumerate(xs): + x_tensor[i, in_max - len(xi):] = torch.tensor(xi, dtype=torch.long, device=device) + + gen = greedy_decode(model, x_tensor, cfg).cpu() + for i, n in enumerate(chunk): + pred = decode_tokens(gen[i].tolist(), cfg) + if is_prime_task: + target = 1 if is_prime_n(n) else 0 + ok = (pred == 1) == (target == 1) # any non-"1" output reads as "composite" + else: + target = next_prime(n, primes) + ok = pred == target + is_easy = (n % 2 == 0) or (n % 5 == 0) # trivial composites (skip-evens / skip-5s) if is_easy: - easy_wrong += 1 + easy_total += 1 + if ok: + correct += 1 + else: + errors.append({"n": n, "target": target, "pred": pred}) + if is_easy: + easy_wrong += 1 total = hi - lo + 1 acc = correct / total if is_prime_task: @@ -124,14 +139,30 @@ def grokking_signature(metrics_path: str) -> dict: @torch.no_grad() def halting_report(model, cfg: Config) -> dict: """RNN halting structure: mean steps at run end + correlation with gap-to-next-prime (H1-H4).""" - primes = sieve_primes(300) + try: + device = next(model.parameters()).device + except StopIteration: + device = torch.device("cpu") + margin = max(100, int(cfg.range_end * 0.05) + 50) + primes = sieve_primes(max(300, cfg.range_end + margin)) gaps, steps = [], [] - for n in range(cfg.range_start, cfg.range_end + 1): - x = torch.tensor(encode_int(n, cfg), dtype=torch.long).unsqueeze(0) - h = model._encode(x) + inputs = list(range(cfg.range_start, cfg.range_end + 1)) + eval_bs = getattr(cfg, "eval_batch_size", 512) + from src.data import pad_inputs + for bi in range(0, len(inputs), eval_bs): + chunk = inputs[bi: bi + eval_bs] + xs = [encode_int(n, cfg) for n in chunk] + in_max = max(len(xi) for xi in xs) + x_tensor = torch.full((len(chunk), in_max), cfg.pad_id, dtype=torch.long, device=device) + for i, xi in enumerate(xs): + x_tensor[i, in_max - len(xi):] = torch.tensor(xi, dtype=torch.long, device=device) + x_padded = pad_inputs(x_tensor, cfg) + h = model._encode(x_padded) _, s = model._run_compute(h) - gaps.append(next_prime(n, primes) - n) - steps.append(float(s.mean())) + s_cpu = s.cpu().tolist() + for i, n in enumerate(chunk): + gaps.append(next_prime(n, primes) - n) + steps.append(float(s_cpu[i])) mean = float(np.mean(steps)) rho = float(np.corrcoef(gaps, steps)[0, 1]) if len(set(gaps)) > 1 else 0.0 lo, hi = cfg.min_steps + 0.5, cfg.max_steps - 0.5 @@ -148,9 +179,12 @@ def halting_report(model, cfg: Config) -> dict: "gaps": gaps, "steps_by_n": steps} -def _load(out_dir: str, ckpt: str, cfg: Config): - model = build_model(cfg) - model.load_state_dict(torch.load(os.path.join(out_dir, ckpt), map_location="cpu")) +def _load(out_dir: str, ckpt: str, cfg: Config, device: torch.device | None = None): + if device is None: + device_str = getattr(cfg, "device", "auto") + device = torch.device("cuda" if (device_str == "cuda" or (device_str == "auto" and torch.cuda.is_available())) else "cpu") + model = build_model(cfg).to(device) + model.load_state_dict(torch.load(os.path.join(out_dir, ckpt), map_location=device)) model.eval() return model |
