Download code/eval_validation_loss.py from laion/Humaneness-Voice-Small: direct link, hf CLI and curl.
- Browser
- Download file 9.87 kB
-
https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/eval_validation_loss.py
- Command line
-
hf download hf://laion/Humaneness-Voice-Small/code/eval_validation_loss.py
-
curl -L -o eval_validation_loss.py https://huggingface.co/laion/Humaneness-Voice-Small/resolve/main/code/eval_validation_loss.py
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() | |