File size: 2,286 Bytes
0662d8e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
"""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())