Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
ChristophSchuhmann's picture
Document architecture, prompts, code, and full run statistics
d911efa verified
Raw History Blame Contribute Delete
4.77 kB
#!/usr/bin/env python3
"""Single-clip Humaneness Voice Small inference; use on a CUDA GPU.
Example (from the model repository root):
python code/infer.py --stage S3 --prompt 'CAPTION: warm, amused narration\nTRANSCRIPT: "Hello there."' \
--text 'Hello there.' --frames 60 --output hello.wav
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import sys
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT / 'code'))
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument('--stage', choices=[f'S{i}' for i in range(1, 11)], default='S3')
parser.add_argument('--prompt', required=True, help='Literal GENERAL/SCRIPT, CAPTION/TRANSCRIPT or TRANSCRIPT text')
parser.add_argument('--text', required=True, help='The exact spoken transcript')
parser.add_argument('--frames', type=int, required=True, help='Frame budget at 12.5 frames/second')
parser.add_argument('--language', choices=('en', 'de'), default='en')
parser.add_argument('--reference-wav', type=Path, help='Optional distinct reference recording')
parser.add_argument('--seed', type=int, default=777)
parser.add_argument('--output', type=Path, required=True)
args = parser.parse_args()
assert args.frames > 0
import numpy as np
import soundfile as sf
import torch
import moss_small
from large_talker import build_fresh
from packing import ScorePacker, generated_audio
if not torch.cuda.is_available():
raise RuntimeError('This inference example requires CUDA')
device = torch.device('cuda:0')
torch.cuda.set_device(device)
moss_small.SFT3 = str(ROOT / 'assets/sft3')
moss_small.QWEN = str(ROOT / 'assets/qwen3')
# Full published model state includes the Qwen3 backbone weights. Do not
# download the separate original pretraining file merely to overwrite it.
moss_small.load_qwen_backbone = lambda model, log=print: None
schema = json.loads((ROOT / 'assets/score_schema.json').read_text())
model, config = build_fresh(schema, log=lambda _: None)
state = torch.load(ROOT / 'checkpoints' / args.stage / 'model_bf16.pt',
map_location='cpu', weights_only=True)
model.load_state_dict(state, strict=True)
model.tie_weights()
model = model.to(device, dtype=torch.bfloat16).eval()
del state
_, _, Processor = moss_small.export_classes()
processor = Processor.from_pretrained(
moss_small.SFT3, codec_path='OpenMOSS-Team/MOSS-Audio-Tokenizer-v2',
codec_weight_dtype='fp32', codec_compute_dtype='bf16')
processor.audio_tokenizer = processor.audio_tokenizer.to(device).eval()
packer = ScorePacker(processor, config, schema)
reference = None
if args.reference_wav:
reference = processor.encode_audios_from_path(args.reference_wav, n_vq=12)[0]
if not 0 < len(reference) <= 37:
raise ValueError('Reference must be 1–37 codec frames (at most 2.96 seconds)')
mode = 'reference' if reference is not None else 'instruction'
record = {'prompt': args.prompt, 'text': args.text, 'frames': args.frames,
'lang': args.language}
example = packer.pack_mode(record, [], mode, reference, generation=True)
batch = packer.collate([example])
ids = batch['input_ids'].to(device)
mask = batch['attention_mask'].to(device)
conditioning = tuple(t.to(device) for t in batch['score_conditioning'])
torch.manual_seed(args.seed)
torch.cuda.manual_seed_all(args.seed)
with (torch.inference_mode(), model.generation_scores(conditioning),
torch.autocast('cuda', dtype=torch.bfloat16)):
result = model.generate(input_ids=ids, attention_mask=mask,
max_new_frames=args.frames + 60, do_sample=True,
audio_temperature=1.0, audio_top_p=0.95, audio_top_k=50,
audio_repetition_penalty=1.0, use_kv_cache=True)
codes = generated_audio(result, config).cpu()
if len(codes) < 2:
raise RuntimeError('Generated fewer than two codec frames')
waveform = processor.decode_audio_codes([codes.to(device)], return_stereo=False)[0]
wave = np.asarray(waveform.float().cpu().numpy()).reshape(-1)
sample_rate = int(processor.model_config.sampling_rate)
assert np.isfinite(wave).all()
assert abs(len(wave) / len(codes) - sample_rate / 12.5) <= 0.01 * sample_rate / 12.5
args.output.parent.mkdir(parents=True, exist_ok=True)
sf.write(args.output, wave, sample_rate, subtype='PCM_16')
print(json.dumps({'output': str(args.output), 'frames': len(codes),
'duration_seconds': len(wave) / sample_rate, 'stage': args.stage}))
if __name__ == '__main__':
main()