OpenWAM-3sys / tools /test_checkpoint_archive.py
qiuly's picture
Add files using upload-large-folder tool
e7aa2d6 verified
Raw History Blame Contribute Delete
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()