matilda-jev-fp4 / kev /continuation.py
yue-maincode's picture
Upload validated MATILDA JEV FP4 model and Decision Index scores
c69aaec verified
Raw History Blame Contribute Delete
2.42 kB
"""Guard an explicitly requested extension while preserving ordinary exact resume checks."""
import hashlib
import math
from pathlib import Path
from kev.train import RESUME_INVARIANT
def validate_extension(meta: dict, config: dict, total_steps: int, hashes: dict, source_run: Path) -> int:
saved = meta['config']
start = int(meta['step'])
if start != meta['total_steps'] or start < 1:
raise ValueError('Only a completed source run may be extended')
if config['epochs'] != saved['epochs'] + 1:
raise ValueError('This extension must add exactly one epoch')
steps_per_epoch = math.ceil(config['train_rows'] / config['effective_batch_size'])
if start != steps_per_epoch * saved['epochs'] or total_steps != start + steps_per_epoch:
raise ValueError('Unexpected epoch or batch schedule')
if meta['examples_seen'] != config['train_rows'] * saved['epochs']:
raise ValueError('Source run did not consume all examples')
if meta['data_sha256'] != hashes or saved['data_sha256'] != hashes:
raise ValueError('Extension data differs from the saved run')
if not 0 < config['lr'] < saved['lr']:
raise ValueError('Extension peak learning rate must be positive and lower')
if not 0 <= config['warmup_fraction'] < 1 or not 0 <= config['min_lr_ratio'] <= 1:
raise ValueError('Invalid extension learning-rate schedule')
if config['world_size'] != saved['world_size']:
raise ValueError('Extension must preserve the FSDP world size')
allowed = {'epochs', 'lr', 'warmup_fraction', 'min_lr_ratio', 'extend_from'}
for key in RESUME_INVARIANT:
if key not in allowed and config[key] != saved[key]:
raise ValueError(f'Extension changes an unrelated setting: {key}')
source_package = source_run.resolve().parents[1] / 'kev/src/kev'
for name, expected in saved['code_sha256'].items():
path = source_package.parents[1] / name if name == 'uv.lock' else source_package / name
actual = hashlib.sha256(path.read_bytes()).hexdigest()
if actual != expected:
raise ValueError(f'Source implementation changed: {name}')
for name in ['model.py', 'evaluate.py', 'types.py', 'uv.lock']:
if config['code_sha256'].get(name) != saved['code_sha256'].get(name):
raise ValueError(f'Extension changes model or evaluation implementation: {name}')
return start