"""Run upstream `train.py` unmodified, against one of our configs. python harness/run_upstream.py configs/rep_s.yaml read 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())