Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
Humaneness-Voice-Small / code /acting_challenge_eval.py
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
5de8603 verified
Raw History Blame Contribute Delete
26 kB
#!/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('<i2').tobytes()
encoder = lameenc.Encoder()
encoder.set_bit_rate(128)
encoder.set_in_sample_rate(sample_rate)
encoder.set_channels(1)
encoder.set_quality(2)
encoded = bytes(encoder.encode(pcm) + encoder.flush())
assert encoded and (encoded.startswith(b'ID3') or encoded[:2] in (b'\xff\xfb', b'\xff\xf3'))
return encoded
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))
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()