Text-to-Speech
English
German
voice-acting
qwen3
moss-audio-tokenizer-v2
audio-generation
File size: 7,108 Bytes
d911efa
 
 
 
88faa07
 
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
88faa07
 
 
d911efa
 
 
 
 
751467a
 
d911efa
 
 
 
751467a
 
88faa07
 
 
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
88faa07
 
 
 
 
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
e770e83
 
 
 
 
 
 
 
751467a
 
 
 
 
 
 
 
 
 
 
 
d911efa
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
88faa07
 
d911efa
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
#!/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 default --reference-wav reference.wav \
      --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=['default', 'S3-Ref15-FullFT'] +
                        [f'S{i}' for i in range(1, 11)], default='default',
                        help='default = S3 Ref15 Full FT; S1–S10 retain the original ladder weights')
    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('--reference-max-seconds', type=float, default=15.0,
                        help='Crop the encoded reference to at most this duration (15 s default; 20 s ceiling)')
    parser.add_argument('--seed', type=int, default=777)
    parser.add_argument('--output', type=Path, required=True)
    args = parser.parse_args()
    assert args.frames > 0
    if not 0 < args.reference_max_seconds <= 20:
        parser.error('--reference-max-seconds must be in (0, 20]')
    if args.stage in ('default', 'S3-Ref15-FullFT') and args.reference_max_seconds > 15:
        print('Warning: Ref15 Full FT trained with at most 14.96 s of reference; '
              'longer inference references are out of distribution', file=sys.stderr)

    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)
    checkpoint_name = 'S3-Ref15-FullFT' if args.stage == 'default' else args.stage
    if checkpoint_name == 'S3-Ref15-FullFT' and args.reference_wav is None:
        print('Warning: the Ref15 default was tuned only with reference audio; '
              'for reference-free inference compare --stage S3', file=sys.stderr)
    state = torch.load(ROOT / 'checkpoints' / checkpoint_name / '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:
        # Torchaudio 2.9 path loading requires optional torchcodec on some
        # installations. SoundFile handles WAV input without that dependency;
        # the original processor still performs codec resampling/encoding.
        reference_wave, reference_rate = sf.read(args.reference_wav, dtype='float32', always_2d=True)
        if not np.isfinite(reference_wave).all():
            raise ValueError('Reference WAV contains non-finite samples')
        reference_tensor = torch.from_numpy(reference_wave.T.copy())
        reference = processor.encode_audios_from_wav([reference_tensor], int(reference_rate), n_vq=12)[0]
        # Historical S1-S10 training used a 37-frame crop, but this is not an
        # architectural cap. Long-reference retraining uses up to 187 frames
        # (14.96 s); for future runs allow a configurable ceiling up to 20 s.
        # Prefer a clean >=5 s recording; shorter references are permitted but
        # should be reported, not mistaken for a full-length conditioning clip.
        max_frames = min(250, int(args.reference_max_seconds / .08))
        reference = reference[:max_frames]
        if not len(reference):
            raise ValueError('Reference codec recording has no frames')
        if len(reference) < 63:
            print(f'Warning: reference is only {len(reference) * .08:.2f}s; prefer at least 5s',
                  file=sys.stderr)
    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, 'resolved_checkpoint': checkpoint_name}))


if __name__ == '__main__':
    main()