File size: 2,255 Bytes
4f03424
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Verify the complete BF16 source against the saved official revision."""
import argparse
import json
import platform
import subprocess
from pathlib import Path
from scripts.common import SOURCE, REVISION, inventory, sha256, write_json

def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--source', type=Path, default=SOURCE)
    ap.add_argument('--output', type=Path, default=Path('artifacts/source-audit.json'))
    ap.add_argument('--manifest', type=Path, help='Published evaluation/source-files.json for a portable rebuild')
    args = ap.parse_args()
    if args.manifest:
        info=json.loads(args.manifest.read_text())
        if info['revision']!=REVISION: raise ValueError('Unexpected source revision')
        rows=info['files']
    else:
        info = json.loads(Path('artifacts/upstream/source-model-info.json').read_text())
        if info['sha'] != REVISION:
            raise ValueError('Unexpected source revision')
        rows = [{'path': x['rfilename'], 'size': x['size'], 'sha256': x['lfs']['sha256']}
                for x in info['siblings'] if x['rfilename'].endswith('.safetensors')]
        rows += [r for r in json.loads(Path('artifacts/upstream/source-config-manifest.json').read_text())
                 if r['path'] not in ('.gitattributes', 'README.md')]
    for row in rows:
        file = args.source / row['path']
        if file.stat().st_size != row['size'] or sha256(file) != row['sha256']:
            raise ValueError(f'Source mismatch: {file}')
        print('Verified', row['path'], flush=True)
    inv = inventory(args.source)
    estimates = {}
    for bits in (4, 6, 8):
        estimates[str(bits)] = sum(v['tensor_bytes'] - v['quantizable_parameters'] * 2 +
                                   v['quantizable_parameters'] * (bits / 8 + 4 / 64)
                                   for v in inv.values())
    write_json(args.output, dict(source=str(args.source), revision=REVISION, files=rows,
                                inventory=inv, estimated_weight_bytes=estimates,
                                platform=platform.platform()))
    print(json.dumps(dict(inventory=inv, estimated_weight_GiB={k:v/2**30 for k,v in estimates.items()}),indent=2))

if __name__ == '__main__':
    main()