#!/usr/bin/env python3 """Generate and score the versioned HVSAC challenge matrix for this Talker.""" from __future__ import annotations import argparse import hashlib import io import json import math import os from pathlib import Path import shutil import subprocess import sys import tarfile 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' NB = Path('/e/data1/datasets/playground/mmlaion/schuhmann1/dramabox') REFERENCE_BANK = SCORE / 'eval/timed_prompts_v3.json' BENCH = Path('/e/home/jusers/schuhmann1/jupiter/m2_600m_tts_training/acting_challenge_eval_v1') REFERENCE_BANK_MANIFEST = BENCH / 'reference_bank.jsonl' 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 encode_mp3_mono(waveform: np.ndarray, sample_rate: int) -> bytes: """Encode one float waveform as 128 kb/s mono MP3 using the local LAME wheel.""" import lameenc samples = np.asarray(waveform, dtype=np.float32).reshape(-1) assert len(samples) and np.isfinite(samples).all() pcm = (np.clip(samples, -1.0, 1.0) * 32767.0).astype(' str: import torch export = checkpoint / 'model_bf16.pt' if export.exists(): assert metadata.get('export_sha256') in (None, sha(export)) 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 modes_for(challenge: dict, surface: int, challenge_index: int) -> tuple[str, ...]: if surface == 1: return ('instruction', 'reference') # paired no-reference/reference C01 add-on if surface == 11: return ('reference',) if surface in (4, 7): return ('instruction',) if surface == 10: return ('reference',) if str(challenge.get('gemma_mode', '')).startswith('REF_') else ('instruction',) if surface in (2, 3, 5, 6, 8, 9): return ('instruction', 'reference')[challenge_index % 2:challenge_index % 2 + 1] raise ValueError(surface) def stage_for(step: int) -> str: bounds = (('S1', 16265), ('S2', 26284), ('S3', 32337), ('S4', 37182), ('S5', 39661), ('S6', 40935), ('S7', 41466), ('S8', 41739), ('S9', 41878), ('S10', 41936)) return next(name for name, end in bounds if step <= end) def write_reference_archive(codec, bank_rows: list[dict], sample_rate: int, output_dir: Path, torch, decode_fn) -> tuple[Path, Path]: """Decode licensed codec references directly into one indexed MP3 TAR.""" output_dir.mkdir(parents=True, exist_ok=True) archive_path = output_dir / 'reference_audio.tar' index_path = output_dir / 'reference_audio.index.json' bank_digest = sha(REFERENCE_BANK_MANIFEST) if archive_path.exists() and index_path.exists(): existing = json.loads(index_path.read_text()) assert existing.get('reference_bank_sha256') == bank_digest assert existing.get('reference_audio_tar_sha256') == sha(archive_path) return archive_path, index_path by_id = {str(row['reference_id']): row for row in bank_rows} tmp_path = archive_path.with_suffix('.tar.tmp') index = {'schema': 'codec-reconstructed-reference-audio-v1', 'reference_bank': 'reference_bank.jsonl', 'reference_bank_sha256': bank_digest, 'reference_audio_tar': 'reference_audio.tar', 'references': {}} with tarfile.open(tmp_path, 'w') as archive: for reference_id, ref in sorted(by_id.items()): codes = torch.as_tensor(ref['codec_codes'], dtype=torch.long, device=codec.audio_tokenizer.device) waveform, valid = decode_fn(codec, codes, sample_rate) assert valid and np.isfinite(waveform).all() and len(waveform) encoded = encode_mp3_mono(waveform, sample_rate) member = f'{reference_id}.mp3' info = tarfile.TarInfo(member) info.size, info.mtime = len(encoded), 0 archive.addfile(info, io.BytesIO(encoded)) index['references'][reference_id] = { 'tar': 'reference_audio.tar', 'member': member, 'codec_reconstructed': True, 'license': ref.get('license'), 'license_evidence': ref.get('license_evidence'), 'source_dataset': ref.get('source_dataset'), 'uid': ref.get('uid'), 'language': ref.get('language'), 'codec_sha256': ref.get('codec_sha256')} tmp_path.replace(archive_path) index['reference_audio_tar_sha256'] = sha(archive_path) tmp_index = index_path.with_suffix('.json.tmp') tmp_index.write_text(json.dumps(index, indent=2, sort_keys=True) + '\n') tmp_index.replace(index_path) return archive_path, index_path def main() -> None: ap = argparse.ArgumentParser() ap.add_argument('--checkpoint', required=True) ap.add_argument('--manifest', default=str(BENCH / 'canary_challenges.jsonl')) ap.add_argument('--reference-bank', default=str(REFERENCE_BANK_MANIFEST)) ap.add_argument('--duration-controls', help='frozen 72-row C01 duration-response manifest') ap.add_argument('--output-root', default=str(SC / 'out/m2_600m_ladder/acting_challenge_eval_v1')) ap.add_argument('--limit', type=int, help='deterministic prefix for canary') ap.add_argument('--dry-run', action='store_true', help='validate manifest/mode/reference inputs only') args = ap.parse_args() checkpoint = Path(args.checkpoint).resolve() metadata = json.loads((checkpoint / 'COMPLETE.json').read_text()) assert int(metadata['step']) > 0 and metadata['world'] == 128 assert (checkpoint / 'training_state.pt').stat().st_size == metadata['state_bytes'] manifest_path = Path(args.manifest).resolve() challenges = [json.loads(line) for line in manifest_path.read_text().splitlines() if line.strip()] assert challenges and len({r['challenge_id'] for r in challenges}) == len(challenges) assert all(list(r['prompts']) == [f'C{i:02d}' for i in range(1, 12)] for r in challenges) assert all(len(r['seeds']) == 3 for r in challenges) assert all(set(str(seed) for seed in r['seeds']) <= set(r.get('prompts_by_seed', {})) for r in challenges), 'seed-specific rendered prompts missing' assert all(set(r['prompts']) == set(r['prompts_by_seed'][str(seed)]) for r in challenges for seed in r['seeds']) if args.limit is not None: assert 1 <= args.limit <= len(challenges) challenges = challenges[:args.limit] challenge_index = {row['challenge_id']: i for i, row in enumerate(challenges)} duration_controls = [] if args.duration_controls: controls_path = Path(args.duration_controls).resolve() control_meta_path = controls_path.with_suffix('.meta.json') assert control_meta_path.is_file(), f'missing frozen control metadata {control_meta_path}' control_meta = json.loads(control_meta_path.read_text()) assert control_meta['challenge_manifest_sha256'] == sha(manifest_path), \ 'duration controls were built from a different benchmark manifest' all_controls = [json.loads(line) for line in controls_path.read_text().splitlines() if line.strip()] assert len(all_controls) == 72 and len({r['task_id'] for r in all_controls}) == 72 duration_controls = [row for row in all_controls if row['challenge_id'] in challenge_index] for control in duration_controls: ci = challenge_index[control['challenge_id']] challenge = challenges[ci] assert control['mode'] in ('instruction', 'reference') assert control['surface_base'] == 'C01' assert control['canonical_sentences'] == challenge['canonical_sentences'] assert control['reference_id'] == challenge['reference_policy']['C11']['reference_id'] assert len(control['duration_targets_s']) == len(challenge['canonical_sentences']) bank_path = Path(args.reference_bank).resolve() bank_rows = [json.loads(line) for line in bank_path.read_text().splitlines() if line.strip()] bank = {str(row.get('reference_id', row.get('ref_uid', row.get('uid')))): row for row in bank_rows} assert bank, 'pinned reference bank is empty' bank_digest = sha(bank_path) task_counts = {'instruction': 0, 'reference': 0} for ci, challenge in enumerate(challenges): policy = challenge.get('reference_policy', {}).get('C11', {}) pinned_id = policy.get('reference_id') or challenge.get('reference_id') assert pinned_id and pinned_id in bank, f"{challenge['challenge_id']} lacks pinned reference asset" reference = bank[pinned_id] codes = reference.get('codec_codes', reference.get('codes', reference.get('ref_codes'))) codec_meta = reference.get('codec_contract', reference.get('codec_metadata', reference.get('codec'))) code_array = np.asarray(codes) assert code_array.ndim == 2 and code_array.shape[1] == 12 and len(code_array) > 0 assert codec_meta, 'reference code/codec contract missing' pinned_uid = str(policy.get('reference_uid', reference.get('uid', reference.get('reference_uid', reference.get('ref_uid'))))) actual_uid = str(reference.get('uid', reference.get('reference_uid', reference.get('ref_uid')))) assert pinned_uid == actual_uid for surface in range(1, 12): for mode in modes_for(challenge, surface, ci): task_counts[mode] += len(challenge['seeds']) for control in duration_controls: task_counts[control['mode']] += 1 if args.dry_run: print(json.dumps({'status': 'PASS', 'challenges': len(challenges), 'tasks': sum(task_counts.values()), 'tasks_by_mode': task_counts, 'challenge_sha256': sha(manifest_path), 'reference_bank_sha256': bank_digest, 'pinned_reference_count': len({ row.get('reference_policy', {}).get('C11', {}).get('reference_id') for row in challenges}), 'duration_control_rows': len(duration_controls), 'duration_control_ids': sorted({r['challenge_id'] for r in duration_controls}), 'selection_role': 'acting benchmark only'}, indent=2)) return import torch import torchaudio import pyarrow as pa import pyarrow.parquet as pq import soundfile as sf from transformers import AutoTokenizer 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)] sys.path[:0] = [str(SC / 'code'), str(NB / 'train2/code/grpo'), str(NB / 'train'), str(SC / 'fastpkgs_sb'), str(NB)] from packing import ScorePacker, generated_audio import moss_small from eval_gen import load_codec_proc, decode from large_talker import build_fresh 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) 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 = json.loads((SCORE / 'score_schema.json').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) codec = load_codec_proc(device) sample_rate = int(codec.model_config.sampling_rate) reference_archive = reference_index = None if rank == 0: reference_archive, reference_index = write_reference_archive( codec, bank_rows, sample_rate, Path(args.output_root), torch, decode) tasks = [(ci, surface, seed, mode, None) for ci, challenge in enumerate(challenges) for surface in range(1, 12) for mode in modes_for(challenge, surface, ci) for seed in challenge['seeds']] tasks.extend((challenge_index[control['challenge_id']], 1, int(control['seed']), control['mode'], control) for control in duration_controls) mine = [task for position, task in enumerate(tasks) if position % world == rank] generated_rows, waves16 = [], [] started = time.monotonic() for ci, surface, seed, mode, control in mine: challenge = challenges[ci] key = f'C{surface:02d}' output_surface = control['surface'] if control else key policy = challenge.get('reference_policy', {}).get('C11', {}) pinned_reference_id = policy.get('reference_id') or challenge.get('reference_id') assert pinned_reference_id, f"{challenge['challenge_id']} has no pinned C11 reference ID" assert pinned_reference_id in bank, f'pinned reference {pinned_reference_id} absent from bank' reference = bank[pinned_reference_id] reference_uid = str(reference.get('uid', reference.get('reference_uid', reference.get('ref_uid')))) reference_codes = reference.get('codec_codes', reference.get('codes', reference.get('ref_codes'))) assert reference_codes, f'pinned reference {pinned_reference_id} has no codec codes' reference_language = str(reference.get('language', reference.get('lang', ''))).lower() expected_language = 'de' if str(challenge['language']).lower().startswith('de') else 'en' assert not reference_language or reference_language.startswith(expected_language) pinned_uid = str(policy.get('reference_uid', reference_uid)) assert pinned_uid == reference_uid assert reference_uid != challenge['challenge_id'] requested_sentence_durations = (control['duration_targets_s'] if control else challenge['duration_targets_s']) target_seconds = float(sum(requested_sentence_durations)) prompt = (control['prompt'] if control else challenge['prompts_by_seed'][str(seed)][key]) meta = {'uid': challenge['challenge_id'] + '_' + output_surface, 'lang': challenge['language'], 'text': challenge['canonical_script'], 'prompt': prompt, 'frames': max(2, int(math.ceil((target_seconds + 5.0) * 12.5)))} torch.manual_seed(int(seed)) torch.cuda.manual_seed_all(int(seed)) example = packer.pack_mode(meta, [], mode, reference_codes if mode == 'reference' else None, generation=True) batch = packer.collate([example]) ids = batch['input_ids'].to(device) mask = batch['attention_mask'].to(device) conditioning = tuple(value.to(device) for value in batch['score_conditioning']) before = time.monotonic() with torch.inference_mode(), model.generation_scores(conditioning), \ torch.autocast('cuda', dtype=torch.bfloat16): generated = model.generate(input_ids=ids, attention_mask=mask, max_new_frames=int(math.ceil(meta['frames'] * 1.6)) + 60, do_sample=True, audio_temperature=1., audio_top_p=.95, audio_top_k=50, audio_repetition_penalty=1., use_kv_cache=True) codes = generated_audio(generated, config).cpu() if len(codes) >= 2: waveform, canary = decode(codec, codes, sample_rate) waveform = np.asarray(waveform, dtype=np.float32) assert canary and np.isfinite(waveform).all() wave16 = torchaudio.functional.resample(torch.from_numpy(waveform), sample_rate, 16000).numpy() else: waveform, wave16, canary = np.zeros(0, np.float32), np.zeros(0, np.float32), None generated_rows.append({'challenge_id': challenge['challenge_id'], 'surface': output_surface, 'mode': mode, 'stage': stage_for(int(metadata['step'])), 'checkpoint_step': int(metadata['step']), 'checkpoint': str(checkpoint), 'seed': int(seed), 'language': challenge['language'], 'script': challenge['canonical_script'], 'prompt': prompt, 'reference_uid': reference_uid if mode == 'reference' else None, 'reference_assignment_uid': pinned_reference_id, 'reference_bank_sha256': bank_digest, 'requested_duration_seconds': target_seconds, 'requested_sentence_durations': [float(x) for x in requested_sentence_durations], 'duration_control_factor': control['duration_factor'] if control else None, 'frames': int(len(codes)), 'duration_seconds': float(len(waveform) / sample_rate), 'canary_ok': canary, 'generation_seconds': time.monotonic() - before, 'waveform': waveform, 'checkpoint_source': checkpoint_source}) waves16.append(wave16) if len(generated_rows) % 10 == 0: print('ACTING_GENERATED', rank, len(generated_rows), flush=True) del model, packer, codec torch.cuda.empty_cache() import asr_parakeet import reward from annotation_scorer import SentenceScorer, NAMES as SCORER_NAMES asr = asr_parakeet.get(device) hypotheses = asr.transcribe(waves16) scorer = SentenceScorer(device) burst_reward = reward.RewardModel(device=str(device), want_asr=False) detections = burst_reward.bursts(waves16) assert len(SCORER_NAMES) == 99 raw_names = ['vn:' + name for name in SCORER_NAMES[:57]] + list(SCORER_NAMES[57:]) assert sum(name.startswith('emo:') for name in raw_names) == 40 assert sum(name.startswith('vn:') for name in raw_names) == 57 raw_rows = [] for wave in waves16: if len(wave) < 400: raw_rows.append(None) continue windows = [wave[start:start + 480000] for start in range(0, len(wave), 480000)] padded = [np.pad(value, (0, max(0, 8000 - len(value)))) for value in windows] raw = scorer.score(padded) durations = np.asarray([len(value) for value in windows], dtype=np.float64) raw_rows.append(np.average(raw, axis=0, weights=durations)) score_rows = [] for row, hypothesis, raw, located in zip(generated_rows, hypotheses, raw_rows, detections): target_challenge = next(item for item in challenges if item['challenge_id'] == row['challenge_id']) target_bursts = target_challenge.get('vocal_bursts', []) requested = [{'label': b['class'], 'sentence_index': int(b['sentence_index']), 'onset_location': b['onset_location'], 'target_duration_seconds': float(b['target_duration_s'])} for b in target_bursts] if raw is None: voicenet, empathy, genuineness, blend = {}, {}, None, None else: voicenet = {name[3:]: float(raw[i]) for i, name in enumerate(raw_names) if name.startswith('vn:')} empathy = {name[4:]: float(raw[i]) for i, name in enumerate(raw_names) if name.startswith('emo:')} genuineness = float(raw[-2]); blend = float(raw[-1]) row_ref = row['script'] import reward as reward_module score_rows.append({**{k: row[k] for k in ('challenge_id','surface','mode','seed','language', 'stage','checkpoint_step','checkpoint','script','prompt','reference_uid', 'reference_assignment_uid','reference_bank_sha256','requested_duration_seconds', 'requested_sentence_durations','duration_control_factor','frames', 'duration_seconds','canary_ok','generation_seconds','checkpoint_source')}, 'hypothesis': hypothesis, 'parakeet_wer': float(reward_module.wer(row_ref, hypothesis)), 'duration_absolute_error_seconds': abs(row['duration_seconds'] - row['requested_duration_seconds']), 'duration_relative_error': (row['duration_seconds'] - row['requested_duration_seconds']) / max(row['requested_duration_seconds'], 1e-6), 'duration_metric_scope': 'total_clip_diagnostic_only', 'sentence_duration_metrics_status': 'pending_forced_alignment', 'duration_response_controls_status': 'pending_72_control_matrix', 'canonical_sentence_targets_json': json.dumps([ {'text': sentence, 'target_duration_seconds': float(target)} for sentence, target in zip(target_challenge['canonical_sentences'], row['requested_sentence_durations'])], ensure_ascii=False, sort_keys=True), 'genuineness_0_6': genuineness, 'vocal_burst_blend_0_10': blend, 'empathic_insight_voice_plus_40_json': json.dumps(empathy, sort_keys=True), 'voicenet_json': json.dumps(voicenet, sort_keys=True), 'requested_bursts_json': json.dumps(requested, sort_keys=True), 'located_bursts_json': json.dumps(located, sort_keys=True), 'burst_realisation_f1': float(reward_module.burst_realisation( [(burst['class'], sum(target_challenge['duration_targets_s'][:max(0, min( int(burst['sentence_index']), len(target_challenge['duration_targets_s'])))] ) + (0.5 * float(target_challenge['duration_targets_s'][int(burst['sentence_index'])]) if burst['onset_location'] == 'during_sentence' and int(burst['sentence_index']) < len(target_challenge['duration_targets_s']) else 0.0), float(burst['target_duration_s'])) for burst in target_bursts], located)), 'audio_seconds': len(waves16[len(score_rows)]) / 16000.0}) out = Path(args.output_root) / checkpoint.name / ('canary' if args.limit else 'full') out.mkdir(parents=True, exist_ok=True) tar_path = out / f'audio.rank{rank:03d}.tar' tmp_tar = tar_path.with_suffix('.tar.tmp') with tarfile.open(tmp_tar, 'w') as archive: for row in generated_rows: if not len(row['waveform']): continue payload = encode_mp3_mono(row['waveform'], sample_rate) info = tarfile.TarInfo(f"{row['challenge_id']}__{row['surface']}__{row['mode']}__seed{row['seed']}.mp3") info.size, info.mtime = len(payload), 0 archive.addfile(info, io.BytesIO(payload)) tmp_tar.replace(tar_path) parquet_path = out / f'scores.rank{rank:03d}.parquet' tmp_parquet = parquet_path.with_suffix('.parquet.tmp') pq.write_table(pa.Table.from_pylist(score_rows), tmp_parquet, compression='zstd') tmp_parquet.replace(parquet_path) proof = {'status': 'PASS', 'rank': rank, 'world': world, 'checkpoint': str(checkpoint), 'checkpoint_step': int(metadata['step']), 'challenge_manifest': str(manifest_path), 'challenge_manifest_sha256': sha(manifest_path), 'reference_bank': str(bank_path), 'reference_bank_sha256': bank_digest, 'reference_bank_manifest': str(bank_path), 'reference_bank_manifest_sha256': bank_digest, 'rows': len(score_rows), 'audio_tar': str(tar_path), 'audio_tar_sha256': sha(tar_path), 'scores': str(parquet_path), 'scores_sha256': sha(parquet_path), 'evaluator_sha256': sha(Path(__file__)), 'elapsed_seconds': time.monotonic() - started, 'browser_audio_format': 'mp3-128k-mono-lameenc', 'reference_audio_tar': str(reference_archive) if reference_archive else None, 'reference_audio_tar_sha256': sha(reference_archive) if reference_archive else None, 'reference_audio_index': str(reference_index) if reference_index else None, 'reference_audio_index_sha256': sha(reference_index) if reference_index else None, 'evaluation_completeness': 'INCOMPLETE_PENDING_SENTENCE_ALIGNMENT_AND_72_DURATION_RESPONSE_CONTROLS', 'publication_ready': False, 'selection_role': 'acting benchmark only; never checkpoint selection'} path = out / f'rank{rank:03d}.PASS.json' tmp = path.with_suffix('.json.tmp'); tmp.write_text(json.dumps(proof, indent=2) + '\n'); tmp.replace(path) print('ACTING_RANK_PASS', rank, len(score_rows), flush=True) if __name__ == '__main__': main()