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()
|