Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
File size: 3,805 Bytes
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
"""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