File size: 1,659 Bytes
4397e12
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Score checkpoints on one common, larger validation set (same windows for every run).
Usage: source env.sh && $TA_PY scripts/rescore.py run_dir [run_dir ...] [--batches 12]"""
import argparse
import json
import os

import numpy as np
import torch

from tiny_agent.checkpoint import load_model
from tiny_agent.data import MixtureLoader
from tiny_agent.model import make_block_mask


@torch.no_grad()
def score(model, batches):
    model.eval()
    out = {}
    for name, bs in batches.items():
        ls = []
        for inp, tgt, doc in bs:
            inp, tgt, doc = inp.to("xpu"), tgt.to("xpu"), doc.to("xpu")
            with torch.autocast("xpu", dtype=torch.bfloat16):
                ls.append(model(inp, doc, make_block_mask(doc, model.cfg.swa_window), tgt).item())
        out[name] = float(np.mean(ls))
    return out


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("runs", nargs="+")
    ap.add_argument("--batches", type=int, default=12)
    a = ap.parse_args()
    batches = MixtureLoader("val", "stable", 2048, 8, stream=False).fixed_batches(a.batches, seed=4321)
    res = {}
    for r in a.runs:
        st = torch.load(os.path.join(r, "ckpt.pt"), map_location="cpu", weights_only=False)
        model = load_model(os.path.join(r, "ckpt.pt"))
        v = score(model, batches)
        res[r] = {"tokens_M": round(st["tokens"] / 1e6), "step": st["step"], "mean_val": round(float(np.mean(list(v.values()))), 4),
                  **{k: round(x, 4) for k, x in sorted(v.items())}}
        print(json.dumps({r: res[r]}), flush=True)
        del model
        torch.xpu.empty_cache()


if __name__ == "__main__":
    main()