File size: 4,687 Bytes
0bbb929
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
66
67
68
69
70
71
72
73
74
75
76
77
78
"""Validation-only selection: score every p2_long export, uniform soups of its last exports, the p1_short final export
and p1-final/p2-final weight averages on the 7,824-row selection split; freeze the lowest-NLL candidate.
The audit partition and the benchmark are not touched here."""
import gc, json, shutil, sys
import numpy as np
import torch
from datasets import load_from_disk
from safetensors.torch import load_file, save_file
from common import *
import engine
from train import exports

SEL = CKPT/'selection'

def soup(paths, dest, weights=None):
    """Weighted (default uniform) average of BF16 exports, computed in FP32, saved as BF16."""
    dest = Path(dest)
    if (dest/'model.safetensors').exists(): return dest
    weights = weights or [1/len(paths)]*len(paths)
    assert abs(sum(weights) - 1) < 1e-6
    acc = None
    for p, w in zip(paths, weights):
        state = load_file(str(Path(p)/'model.safetensors'))
        if acc is None: acc = {k: v.float()*w for k, v in state.items()}
        else:
            assert acc.keys() == state.keys()
            for k, v in state.items(): acc[k] += v.float()*w
    stage = dest.with_name(dest.name+'.staging')
    if stage.exists(): shutil.rmtree(stage)
    shutil.copytree(Path(paths[0]), stage, ignore=shutil.ignore_patterns('model.safetensors', '*.npz', 'training_progress.json', 'soup.json'))
    save_file({k: v.to(torch.bfloat16).contiguous() for k, v in acc.items()}, str(stage/'model.safetensors'), metadata={'format': 'pt'})
    atomic_json(stage/'soup.json', {'members': [str(p) for p in paths], 'weights': weights})
    stage.rename(dest)
    return dest

def evaluate(path, name, val, indices):
    out = SEL/(name+'.npz')
    model = engine.load_model(path)
    logits = engine.predict(model, val, indices, 'selection-'+name, out)
    del model; gc.collect(); torch.cuda.empty_cache()
    r = engine.report(val, indices, logits)
    return {'name': name, 'path': str(path), 'sha256': sha(Path(path)/'model.safetensors'),
            'nll': r['overall']['nll'], 'balanced_accuracy': r['overall']['balanced_accuracy'], 'accuracy': r['overall']['accuracy'],
            'false_approve_count': r['overall']['false_approve_count'], 'false_deny_count': r['overall']['false_deny_count'],
            'long_accuracy': r['slices']['length']['16k-64k']['accuracy'], 'auroc': r['overall']['auroc'], 'report': r}

def main(run='p2_long'):
    if (SEL/'summary.json').exists() and '--force' not in sys.argv:
        print(json.dumps(json.loads((SEL/'summary.json').read_text())['frozen'], indent=1)); return
    SEL.mkdir(parents=True, exist_ok=True)
    val = load_from_disk(str(DATA/'validation')); indices = np.load(DATA/'validation_partitions.npz')['selection']
    assert (CKPT/run/'complete.json').exists()
    ex = exports(run); results = {}
    for p in ex: results[f'{run}-{p.name}'] = evaluate(p, f'{run}-{p.name}', val, indices)
    p1 = exports('p1_short')[-1]
    results[f'p1_short-{p1.name}'] = evaluate(p1, f'p1_short-{p1.name}', val, indices)
    for k in (2, 3, 4):
        if len(ex) >= k:
            name = f'soup-{run}-last{k}'
            results[name] = evaluate(soup(ex[-k:], SEL/name), name, val, indices)
    # Added before any benchmark evaluation (p2 raised validation NLL on short inputs): p1-final/p2-final averages.
    for w in (0.3, 0.5, 0.7):
        name = f'soup-p1final{w:.1f}+p2final{1-w:.1f}'
        results[name] = evaluate(soup([p1, ex[-1]], SEL/name, [w, round(1-w, 1)]), name, val, indices)
    eligible = list(results)
    ranked = sorted(eligible, key=lambda n: results[n]['nll'])
    frozen = results[ranked[0]]
    for r in sorted(results.values(), key=lambda r: r['nll']):
        print(f"{'*' if r['name'] == frozen['name'] else ' '}{r['name']:32s} nll={r['nll']:.5f} bacc={r['balanced_accuracy']:.5f} acc={r['accuracy']:.5f} "
              f"FA={r['false_approve_count']:3d} FD={r['false_deny_count']:3d} long={r['long_accuracy']:.4f}", flush=True)
    atomic_json(SEL/'summary.json', {'frozen': {k: v for k, v in frozen.items() if k != 'report'}, 'ranked': ranked,
                'criterion': 'lowest NLL on the 7,824-row validation selection split among p2_long exports, last-2/3/4 uniform p2 soups, the p1_short final export and p1-final/p2-final weight averages (0.3/0.5/0.7); frozen before any benchmark or audit evaluation',
                'candidates': [{k: v for k, v in r.items() if k != 'report'} for r in sorted(results.values(), key=lambda r: r['nll'])]})
    for r in results.values(): atomic_json(SEL/(r['name']+'.report.json'), r['report'])
    event('frozen', **{k: v for k, v in frozen.items() if k != 'report'})

if __name__ == '__main__': main()