Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /eval_validation_loss.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
5de8603 verified
Raw History Blame Contribute Delete
9.87 kB
#!/usr/bin/env python3
"""Teacher-forced held-out loss for continuous large-Talker checkpoints.
The fixed 50-row evaluation set is sourced from the benchmark's explicitly
held-out split. It is independent of the stage manifests in the production
plan. Each rank writes one Parquet shard and a small proof record.
"""
from __future__ import annotations
import argparse
import hashlib
import json
import math
import os
from pathlib import Path
import sys
import time
import numpy as np
SC = Path('/e/scratch/reformo/schuhmann1_moss')
HERE = Path(__file__).resolve().parent
V3 = Path('/e/home/jusers/schuhmann1/jupiter/m2_600m_tts_training/cascade_v3_timed')
M2 = SC / 'code/m2_20k_score'
SMALL = SC / 'code/small_tts'
SCORE = SC / 'out/m2_20k_score'
DEFAULT_PROMPTS = SCORE / 'eval/timed_prompts_v3.json'
DEFAULT_OUT = SC / 'out/m2_600m_ladder/caption_curriculum_v1/validation_loss'
def sha(path: Path) -> str:
digest = hashlib.sha256()
with path.open('rb') as stream:
for block in iter(lambda: stream.read(8 << 20), b''):
digest.update(block)
return digest.hexdigest()
def audit_heldout(prompts: list[dict], plan: dict) -> None:
assert len(prompts) == 50, f'expected fixed held-out set of 50 rows, got {len(prompts)}'
assert len({row['uid'] for row in prompts}) == len(prompts), 'duplicate held-out UID'
assert all(row.get('split') == 'heldout' for row in prompts), 'non-heldout prompt present'
# The ladder release contract defines tier 8 as a UID-disjoint held-out
# partition. Stage manifests contain only tier manifests, never heldout.
assert all('heldout' not in str(stage['manifest']).lower() for stage in plan['stages'])
assert all(row.get('codes') and row.get('ref_codes') for row in prompts), \
'held-out row lacks target or validated reference codes'
def load_checkpoint(model, checkpoint: Path, metadata: dict) -> str:
import torch
export = checkpoint / 'model_bf16.pt'
if export.exists():
assert metadata.get('export_sha256') in (None, sha(export)), 'export SHA mismatch'
state = torch.load(export, map_location='cpu', weights_only=True)
source = 'model_bf16.pt'
else:
state_path = checkpoint / 'training_state.pt'
assert state_path.stat().st_size == metadata['state_bytes']
assert sha(state_path) == metadata['state_sha256']
full = torch.load(state_path, map_location='cpu', weights_only=False)
state = {key[5:]: value for key, value in full['model'].items()
if key.startswith('base.')}
assert state and len(state) == len(full['model'])
source = 'training_state.pt:model'
del full
missing, unexpected = model.load_state_dict(state, strict=True)
assert not missing and not unexpected
model.tie_weights()
del state
return source
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument('--checkpoint', required=True)
parser.add_argument('--prompts', default=str(DEFAULT_PROMPTS))
parser.add_argument('--output-root', default=str(DEFAULT_OUT))
parser.add_argument('--limit', type=int, help='deterministic prefix for a canary')
args = parser.parse_args()
checkpoint = Path(args.checkpoint).resolve()
metadata_path = checkpoint / 'COMPLETE.json'
metadata = json.loads(metadata_path.read_text())
assert int(metadata['step']) > 0 and metadata['world'] == 128
assert (checkpoint / 'training_state.pt').stat().st_size == metadata['state_bytes']
plan_path = HERE / 'plan_caption_production_32nodes.json'
plan = json.loads(plan_path.read_text())
prompts_path = Path(args.prompts).resolve()
prompts = json.loads(prompts_path.read_text())
audit_heldout(prompts, plan)
if args.limit is not None:
assert 1 <= args.limit <= len(prompts)
prompts = prompts[:args.limit]
rank = int(os.environ.get('SLURM_PROCID', '0'))
world = int(os.environ.get('SLURM_NTASKS', '1'))
local = int(os.environ.get('SLURM_LOCALID', '0'))
assert world in (1, 4), f'evaluation expects one or four ranks, got {world}'
for candidate in (str(HERE), str(V3), str(M2), str(SMALL)):
while candidate in sys.path:
sys.path.remove(candidate)
sys.path[:0] = [str(HERE), str(V3), str(M2), str(SMALL)]
import torch
from transformers import AutoTokenizer
from packing import ScorePacker
import moss_small
import va_loss
from large_talker import build_fresh
device_index = 0 if torch.cuda.device_count() == 1 else local
torch.cuda.set_device(device_index)
device = torch.device('cuda', device_index)
torch.set_num_threads(4)
schema_path = SCORE / 'score_schema.json'
schema = json.loads(schema_path.read_text())
model, config = build_fresh(schema, log=lambda _: None)
checkpoint_source = load_checkpoint(model, checkpoint, metadata)
model = model.to(device, dtype=torch.bfloat16).eval()
_, _, Processor = moss_small.export_classes()
processor = Processor(tokenizer=AutoTokenizer.from_pretrained(
moss_small.SFT3, trust_remote_code=True, local_files_only=True),
audio_tokenizer=None, model_config=config)
packer = ScorePacker(processor, config, schema)
tasks = [(index, mode) for index in range(len(prompts))
for mode in ('instruction', 'reference')]
mine = [task for position, task in enumerate(tasks) if position % world == rank]
started = time.monotonic()
rows = []
per_channel_numerator = np.zeros(13, dtype=np.float64)
per_channel_denominator = np.zeros(13, dtype=np.float64)
for index, mode in mine:
prompt = prompts[index]
reference = prompt['ref_codes'] if mode == 'reference' else None
example = packer.pack_mode(prompt, prompt['codes'], mode, reference, generation=False)
batch = packer.collate([example])
inputs = {
'input_ids': batch['input_ids'].to(device),
'attention_mask': batch['attention_mask'].to(device),
'labels': batch['labels'].to(device),
'score_conditioning': tuple(value.to(device) for value in batch['score_conditioning']),
}
with torch.inference_mode(), torch.autocast('cuda', dtype=torch.bfloat16):
hidden = model(input_ids=inputs['input_ids'], attention_mask=inputs['attention_mask'],
score_conditioning=inputs['score_conditioning'],
use_cache=False).last_hidden_state
loss, per = va_loss.compute_supervised_loss_from_hidden(
model, global_hidden_states=hidden, labels=inputs['labels'],
channelwise_loss_weight=[1.0] + [32.0 / 12.0] * 12,
return_per_channel=True)
frame_count = int(batch['labels'][:, :, 1].ge(0).sum().item())
sample_count = 1
assert frame_count == len(prompt['codes'])
per_values = per.detach().float().cpu().numpy().reshape(-1)
assert len(per_values) == 13 and np.isfinite(per_values).all()
counts = np.asarray([frame_count + sample_count] + [frame_count] * 12,
dtype=np.float64)
per_channel_numerator += per_values * counts
per_channel_denominator += counts
rows.append({
'uid': prompt['uid'], 'split': prompt['split'], 'source': prompt['source'],
'lang': prompt['lang'], 'mode': mode, 'checkpoint_step': int(metadata['step']),
'frames': frame_count, 'samples': sample_count,
'loss': float(loss.detach().float().cpu()),
'per_channel_loss_json': json.dumps(per_values.tolist()),
})
del hidden, batch, example
if len(rows) % 10 == 0:
print('VALIDATION_ROW', rank, len(rows), flush=True)
import pyarrow as pa
import pyarrow.parquet as pq
destination = Path(args.output_root) / checkpoint.name / ('canary' if args.limit else 'full')
destination.mkdir(parents=True, exist_ok=True)
parquet_path = destination / f'loss.rank{rank:03d}.parquet'
temporary = parquet_path.with_suffix('.parquet.tmp')
pq.write_table(pa.Table.from_pylist(rows), temporary, compression='zstd')
temporary.replace(parquet_path)
channels = np.divide(per_channel_numerator, per_channel_denominator,
out=np.zeros_like(per_channel_numerator),
where=per_channel_denominator > 0)
weights = np.asarray([1.0] + [32.0 / 12.0] * 12)
proof = {
'status': 'PASS', 'rank': rank, 'world': world,
'checkpoint': str(checkpoint), 'checkpoint_step': int(metadata['step']),
'checkpoint_contract': metadata['contract'],
'checkpoint_source': checkpoint_source,
'heldout': True, 'heldout_rows_in_input': len(prompts), 'task_rows': len(rows),
'prompts_path': str(prompts_path), 'prompts_sha256': sha(prompts_path),
'schema_sha256': sha(schema_path), 'production_plan_sha256': sha(plan_path),
'evaluator_sha256': sha(Path(__file__)),
'per_channel_loss': channels.tolist(),
'per_channel_numerator': per_channel_numerator.tolist(),
'per_channel_denominator': per_channel_denominator.tolist(),
'validation_loss': float(np.dot(channels, weights) / weights.sum()),
'parquet': str(parquet_path), 'parquet_sha256': sha(parquet_path),
'elapsed_seconds': time.monotonic() - started,
'slurm_job_id': os.environ.get('SLURM_JOB_ID'),
}
proof_path = destination / f'rank{rank:03d}.PASS.json'
temporary_proof = proof_path.with_suffix('.json.tmp')
temporary_proof.write_text(json.dumps(proof, indent=2) + '\n')
temporary_proof.replace(proof_path)
print('VALIDATION_RANK_PASS', rank, len(rows), proof['validation_loss'], flush=True)
if __name__ == '__main__':
main()