Download reproducibility/stage_checkpoint.py from ProCreations/Agnes-3.0-Flash-NVFP4: direct link, hf CLI and curl.
- Browser
- Download file 2.41 kB
-
https://huggingface.co/ProCreations/Agnes-3.0-Flash-NVFP4/resolve/fc2119a0164ef0e766f4889754d3d16440f3478a/reproducibility/stage_checkpoint.py
- Command line
-
hf download hf://ProCreations/Agnes-3.0-Flash-NVFP4@fc2119a0164ef0e766f4889754d3d16440f3478a/reproducibility/stage_checkpoint.py
-
curl -L -o stage_checkpoint.py https://huggingface.co/ProCreations/Agnes-3.0-Flash-NVFP4/resolve/fc2119a0164ef0e766f4889754d3d16440f3478a/reproducibility/stage_checkpoint.py
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) | |