#!/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()