Download tools/test_checkpoint_archive.py from qiuly/OpenWAM-3sys: direct link, hf CLI and curl.
- Browser
- Download file 6.59 kB
-
https://huggingface.co/qiuly/OpenWAM-3sys/resolve/main/tools/test_checkpoint_archive.py
- Command line
-
hf download hf://qiuly/OpenWAM-3sys/tools/test_checkpoint_archive.py
-
curl -L -o test_checkpoint_archive.py https://huggingface.co/qiuly/OpenWAM-3sys/resolve/main/tools/test_checkpoint_archive.py
6.59 kB
| import concurrent.futures | |
| import json | |
| from pathlib import Path | |
| import struct | |
| import tempfile | |
| import unittest | |
| from checkpoint_archive import Archive, checked_path, sha256_file | |
| def source(root, run, reason=b'12345678', step=10): | |
| folder = root / ('source-' + run) | |
| folder.mkdir() | |
| tensors = { | |
| 'action_backbone.test': {'dtype': 'F32', 'shape': [2], 'data_offsets': [0, 8]}, | |
| 'video_backbone.reason1.test': {'dtype': 'BF16', 'shape': [4], 'data_offsets': [8, 16]}, | |
| 'vlm_expert.vlm.visual.test': {'dtype': 'BF16', 'shape': [2], 'data_offsets': [16, 20]}, | |
| } | |
| # Non-canonical whitespace/metadata must survive byte-identical restoration. | |
| tensors['__metadata__'] = {'description': 'original header preserved'} | |
| raw = json.dumps(tensors, indent=1).encode() | |
| path = folder / f'checkpoint_step_{step}.safetensors' | |
| path.write_bytes(struct.pack('<Q', len(raw)) + raw + b'abcdefgh' + reason + b'WXYZ') | |
| (folder / 'config.yaml').write_text('model: {}\n') | |
| (folder / 'normalization_stats.npy').write_bytes(b'stats-' + run.encode()) | |
| (folder / 'reason1').mkdir() | |
| (folder / 'reason1/config.json').write_text('{}') | |
| return path | |
| class ArchiveTest(unittest.TestCase): | |
| def setUp(self): | |
| self.temp = tempfile.TemporaryDirectory() | |
| self.base = Path(self.temp.name) | |
| self.a = Archive(self.base / 'archive') | |
| self.a.init() | |
| def tearDown(self): | |
| self.temp.cleanup() | |
| def add(self, run, **kwargs): | |
| src = source(self.base, run, **kwargs) | |
| return src, self.a.add_policy(run=run, name=run + '_test', stage='sft', source=src) | |
| def test_lossless_relocation_and_real_copies(self): | |
| src, e = self.add('run1') | |
| self.a.verify(reconstruct=True) | |
| relocated = self.base / 'moved' | |
| self.a.root.rename(relocated) | |
| new = Archive(relocated) | |
| output = new.materialize('run1', self.base / 'restored') | |
| self.assertEqual(src.read_bytes(), (output / src.name).read_bytes()) | |
| self.assertEqual((src.parent / 'config.yaml').read_bytes(), (output / 'config.yaml').read_bytes()) | |
| self.assertNotEqual(src.stat().st_ino, (output / src.name).stat().st_ino) | |
| self.assertFalse(any(p.is_symlink() for p in relocated.rglob('*'))) | |
| def test_shared_exact_match_and_changed_component(self): | |
| self.add('run1') | |
| self.add('run2') | |
| self.assertEqual(len(list((self.a.root / 'shared/reason1').glob('*/model.safetensors'))), 1) | |
| self.add('run3', reason=b'changed!') | |
| self.assertEqual(len(list((self.a.root / 'shared/reason1').glob('*/model.safetensors'))), 2) | |
| self.a.verify(reconstruct=True) | |
| def test_incremental_idempotence_and_conflict(self): | |
| src, before = self.add('run40') | |
| self.a.add_policy(run='run40', name='run40_test', stage='sft', source=src) | |
| self.add('run41') | |
| self.assertEqual(before, self.a.resolve('run40')) | |
| changed = source(self.base, 'different', reason=b'new data') | |
| with self.assertRaises(FileExistsError): | |
| self.a.add_policy(run='run40', name='run40_test', stage='sft', source=changed) | |
| (src.parent / 'config.yaml').write_text('changed: true') | |
| with self.assertRaises(FileExistsError): | |
| self.a.add_policy(run='run40', name='run40_test', stage='sft', source=src) | |
| def test_corruption_never_publishes_materialized_output(self): | |
| src, e = self.add('run1') | |
| original = sha256_file(src) | |
| record = self.a.index(e)['files'][0] | |
| dest = self.a.root / record['path'] | |
| raw = bytearray(dest.read_bytes()); raw[-1] ^= 1; dest.write_bytes(raw) | |
| with self.assertRaises(ValueError): | |
| self.a.materialize('run1', self.base / 'restore') | |
| self.assertFalse((self.base / 'restore').exists()) | |
| self.assertEqual(sha256_file(src), original) | |
| def test_incomplete_cache_marker_cannot_hide_missing_weights(self): | |
| src, _ = self.add('run1') | |
| out = self.a.materialize('run1', self.base / 'restore') | |
| marker = out / '.archive-materialized.json' | |
| data = json.loads(marker.read_text()) | |
| data['files'] = [] | |
| marker.write_text(json.dumps(data)) | |
| (out / src.name).unlink() | |
| with self.assertRaises(ValueError): | |
| self.a.materialize('run1', out) | |
| def test_added_metadata_is_not_silently_ignored_on_reimport(self): | |
| src, _ = self.add('run1') | |
| (src.parent / 'new-config.json').write_text('{}') | |
| with self.assertRaises(FileExistsError): | |
| self.a.add_policy(run='run1', name='run1_test', stage='sft', source=src) | |
| def test_missing_dependency_never_enters_catalog(self): | |
| src = source(self.base, 'run1') | |
| with self.assertRaises(ValueError): | |
| self.a.add_policy(run='run1', name='run1_test', stage='sft', source=src, dependencies={'svae': 'svae/run99'}) | |
| self.assertEqual(self.a.catalog()['entries'], {}) | |
| self.assertTrue(src.exists()) | |
| def test_portable_model_and_dependency(self): | |
| src = self.base / 'svae'; src.mkdir(); (src / 'svae.pt').write_bytes(b'complete model') | |
| self.a.add_model(key='svae/run1', name='run1_svae', source=src, files=['svae.pt']) | |
| ckpt = source(self.base, 'run1') | |
| self.a.add_policy(run='run1', name='run1_test', stage='sft', source=ckpt, dependencies={'svae': 'svae/run1'}) | |
| self.a.verify(reconstruct=True) | |
| out = self.a.materialize('svae/run1', self.base / 'restored-model') | |
| self.assertEqual((out / 'svae.pt').read_bytes(), b'complete model') | |
| def test_path_and_symlink_guards(self): | |
| with self.assertRaises(ValueError): | |
| checked_path(self.a.root, '../other') | |
| with self.assertRaises(ValueError): | |
| checked_path(self.a.root, '/absolute') | |
| (self.a.root / 'link').symlink_to(self.base) | |
| with self.assertRaises(ValueError): | |
| checked_path(self.a.root, 'link/anything') | |
| def test_concurrent_append_preserves_both_catalog_entries(self): | |
| sources = [source(self.base, 'run1'), source(self.base, 'run2')] | |
| def add(pair): | |
| i, src = pair | |
| return Archive(self.a.root).add_policy(run=f'run{i}', name=f'run{i}_test', stage='sft', source=src) | |
| with concurrent.futures.ThreadPoolExecutor(max_workers=2) as pool: | |
| list(pool.map(add, enumerate(sources, 1))) | |
| self.assertEqual(set(self.a.catalog()['entries']), {'sft/run01', 'sft/run02'}) | |
| self.a.verify(reconstruct=True) | |
| if __name__ == '__main__': | |
| unittest.main() | |