mini-agi-replication / code /harness /run_upstream.py
dreddnafious's picture
Revision 2: step density matched; headline replicates
0662d8e verified
Raw History Blame Contribute Delete
2.29 kB
"""Run upstream `train.py` unmodified, against one of our configs.
python harness/run_upstream.py configs/rep_s.yaml read <train.py args...>
The only behavioural change is logging-only and opt-in: SAMPLE_GEN_EVERY=N
makes train.sample_now generate text on every Nth call and return the previous
samples otherwise. Held-out evaluation - which drives the learning-rate
controller - still runs on every call; only the greedy text samples, which
cost more than the evaluation at this scale and feed nothing, are thinned.
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from upstream import use_config # noqa: E402
def _patch_grad_snr():
"""Upstream GradSNR crashes when a parameter's grad is None on some steps
but not others: a depth-1 step leaves the halting head without a gradient,
and the running mean then changes length. Upstream's depth (Poisson 12.8)
almost never samples 1; ours (4.3) does ~1.4% of steps. The meter is
reporting-only, so a missing grad is counted as zeros."""
import torch
from minagi.optim import GradSNR
@torch.no_grad()
def observe(self, params):
gs = [p.grad if p.grad is not None else torch.zeros_like(p) for p in params]
if all(p.grad is None for p in params):
return None
flat = torch.cat([g.detach().float().reshape(-1) for g in gs])
self.m = flat.clone() if self.m is None else \
self.m.mul_(self.beta).add_(flat, alpha=1 - self.beta)
s = float((flat * flat).sum())
self.sq = s if self.n == 0 else self.beta * self.sq + (1 - self.beta) * s
self.n += 1
return self.ratio()
GradSNR.observe = observe
def main():
use_config(sys.argv[1])
_patch_grad_snr()
import train
every = int(os.environ.get("SAMPLE_GEN_EVERY", "1"))
if every > 1:
real, state = train.sample_now, {"n": 0, "last": None}
def thinned(*a, **k):
if state["last"] is None or state["n"] % every == 0:
state["last"] = real(*a, **k)
state["n"] += 1
return state["last"]
train.sample_now = thinned
sys.argv = ["train.py"] + sys.argv[2:]
return train.main()
if __name__ == "__main__":
sys.exit(main())