File size: 2,420 Bytes
c69aaec
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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