summaryrefslogtreecommitdiff
path: root/src/eval.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/eval.py')
-rw-r--r--src/eval.py90
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