Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
3.81 kB
"""Atomic full-state DDP checkpoints with per-rank RNG and pinned contracts."""
from datetime import datetime
import json
import os
from pathlib import Path
import random
import numpy as np
import torch
import torch.distributed as dist
from continuation_control import sha
def ranks():
return (dist.get_rank(), dist.get_world_size()) if dist.is_initialized() else (0, 1)
def barrier():
if dist.is_initialized():
dist.barrier()
def unwrap(model):
return model.module if hasattr(model, 'module') else model
def save(root, model, optimizer, scheduler, step, contract, export=False):
rank, world = ranks()
root = Path(root)
rng = {'torch': torch.get_rng_state(), 'cuda': torch.cuda.get_rng_state(),
'numpy': np.random.get_state(), 'python': random.getstate()}
all_rng = [None] * world if rank == 0 else None
if world > 1:
dist.gather_object(rng, all_rng, dst=0)
else:
all_rng = [rng]
directory = root / f'step-{step:08d}'
if rank == 0:
temporary = root / f'.step-{step:08d}-{os.environ.get("SLURM_JOB_ID", "local")}.partial'
temporary.mkdir(parents=True, exist_ok=True)
torch.save({'model': unwrap(model).state_dict(), 'optimizer': optimizer.state_dict(),
'scheduler': scheduler.state_dict(), 'rng': all_rng, 'step': step, 'world': world,
'contract': contract}, temporary / 'training_state.pt')
metadata = {'step': step, 'world': world, 'contract': contract,
'created_at': datetime.now().astimezone().isoformat(),
'state_bytes': (temporary / 'training_state.pt').stat().st_size,
'state_sha256': sha(temporary / 'training_state.pt')}
if export:
base = unwrap(model).base
cache, state = {}, {}
for name, value in base.state_dict().items():
key = (value.data_ptr(), tuple(value.shape), value.dtype)
if key not in cache:
cache[key] = value.detach().to(dtype=torch.bfloat16, device='cpu')
state[name] = cache[key]
torch.save(state, temporary / 'model_bf16.pt')
metadata['export_bytes'] = (temporary / 'model_bf16.pt').stat().st_size
metadata['export_sha256'] = sha(temporary / 'model_bf16.pt')
(temporary / 'COMPLETE.json').write_text(json.dumps(metadata, indent=2) + '\n')
if directory.exists():
raise RuntimeError(f'Refusing to overwrite existing checkpoint {directory}')
temporary.replace(directory)
pointer = root / 'current.json.tmp'
pointer.write_text(json.dumps({'directory': directory.name, 'step': step}) + '\n')
pointer.replace(root / 'current.json')
barrier()
return directory
def load(directory, model, optimizer, scheduler, contract):
directory = Path(directory)
rank, world = ranks()
metadata = json.loads((directory / 'COMPLETE.json').read_text())
assert metadata['contract'] == contract and metadata['world'] == world
assert (directory / 'training_state.pt').stat().st_size == metadata['state_bytes']
state = torch.load(directory / 'training_state.pt', map_location='cpu', weights_only=False)
assert state['contract'] == contract and state['world'] == world and len(state['rng']) == world
unwrap(model).load_state_dict(state['model'], strict=True)
optimizer.load_state_dict(state['optimizer']); scheduler.load_state_dict(state['scheduler'])
rng = state['rng'][rank]
torch.set_rng_state(rng['torch']); torch.cuda.set_rng_state(rng['cuda'])
np.random.set_state(rng['numpy']); random.setstate(rng['python'])
step = int(state['step'])
assert step==metadata['step']
del state
barrier()
return step