Agnes-3.0-Flash-NVFP4 / reproducibility /stage_checkpoint.py
ProCreations's picture
Pin matched Agnes native validation suite
63f6cfd verified
Raw History Blame
2.41 kB
"""Copy mounted inference files sequentially and verify Hub content hashes."""
import hashlib
import json
from pathlib import Path
PROTOCOL = 'sequential-local-copy-hub-hash-v1'
def stage_checkpoint(mount, destination, manifest_path):
mount, destination, manifest_path = map(Path, (mount, destination, manifest_path))
manifest_bytes = manifest_path.read_bytes()
manifest = json.loads(manifest_bytes)
destination.mkdir(parents=True, exist_ok=True)
total = 0
for item in manifest['files']:
name = item['name']
assert not Path(name).is_absolute() and '..' not in Path(name).parts, 'Unsafe checkpoint relative path'
(destination/name).parent.mkdir(parents=True,exist_ok=True)
source = mount/name
assert source.stat().st_size == item['size'], 'Mounted file size differs: '+name
print('SERVER_MODEL_STAGE_BEGIN', name, item['size'], flush=True)
if item['hash_algorithm'] == 'sha256':
digest = hashlib.sha256()
else:
assert item['hash_algorithm'] == 'git-sha1'
digest = hashlib.sha1()
digest.update(('blob '+str(item['size'])+'\0').encode())
temporary = destination/(name+'.partial')
count = 0; next_progress = 512*1024*1024
try:
with source.open('rb') as reader, temporary.open('wb') as writer:
while block := reader.read(8*1024*1024):
writer.write(block); digest.update(block); count += len(block)
if count >= next_progress:
print('SERVER_MODEL_STAGE_PROGRESS', name, count, item['size'], flush=True)
next_progress += 512*1024*1024
assert count == item['size'], 'Incomplete file copy: '+name
assert digest.hexdigest() == item['hash'], 'Content hash differs: '+name
temporary.replace(destination/name)
except Exception:
temporary.unlink(missing_ok=True)
raise
total += count
print('SERVER_MODEL_FILE_VERIFIED', name, count, flush=True)
print('SERVER_MODEL_STAGE_COMPLETE', len(manifest['files']), total, flush=True)
return dict(protocol=PROTOCOL, manifest_sha256=hashlib.sha256(manifest_bytes).hexdigest(),
source_files_revision=manifest['source_files_revision'], verified_files=len(manifest['files']),
verified_bytes=total)