"""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