File size: 2,346 Bytes
d911efa | 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 | """Pure control rules shared by the real trainer and its recovery verifier."""
import hashlib
import json
import math
from pathlib import Path
def sha(path):
result=hashlib.sha256()
with Path(path).open('rb') as stream:
for block in iter(lambda:stream.read(8<<20),b''):result.update(block)
return result.hexdigest()
def complete_checkpoint(directory,contract=None,world=None,content=False):
directory=Path(directory);meta=json.loads((directory/'COMPLETE.json').read_text())
assert meta['step']>=1
if contract is not None:assert meta['contract']==contract
if world is not None:assert meta['world']==world
assert (directory/'training_state.pt').stat().st_size==meta['state_bytes']>0
if content:assert sha(directory/'training_state.pt')==meta['state_sha256']
if 'export_bytes' in meta:
assert (directory/'model_bf16.pt').stat().st_size==meta['export_bytes']>0
if content:assert sha(directory/'model_bf16.pt')==meta['export_sha256']
return meta
def latest_checkpoint(root,contract,world):
"""A crash after directory rename but before pointer update is recoverable."""
found=[]
for directory in Path(root).glob('step-*'):
# A named committed directory with invalid content is a real error;
# never silently rewind a possibly corrupted training history.
meta=complete_checkpoint(directory,contract,world)
assert directory.name==f'step-{meta["step"]:08d}'
found.append((meta['step'],directory))
return max(found)[1] if found else None
def epoch_plan(samples,world,per_gpu):
assert samples>0 and world>0 and per_gpu>0
batch=world*per_gpu;steps=math.ceil(samples/batch);tail=samples-(steps-1)*batch
assert tail>=world,'Every final rank must receive a real target; no silent tail drop/padding'
quarters=sorted({math.ceil(steps*q/4) for q in (1,2,3,4)})
return {'samples':samples,'world':world,'global_batch':batch,'steps':steps,'tail_samples':tail,
'quarter_steps':quarters,'epochs':1,'presentations_per_target':1,'dropped_targets':0,'duplicated_targets':0}
def contract(plan_path,plan,world):
return {'plan_sha256':sha(plan_path),'schema_sha256':plan['schema_sha256'],'manifest_sha256':plan['manifest_sha256'],
'world':world,'objective':'global_token_mean_v1','seed':plan['seed']}
|