#!/usr/bin/env python3 """Inference entry point for Gooya Bozorg v1.5.""" from __future__ import annotations import argparse from pathlib import Path import soundfile as sf from safetensors.torch import load_file import torch try: from chatterbox.models.s3gen import S3Gen from chatterbox.models.t3 import T3 from chatterbox.models.t3.modules.t3_config import T3Config from chatterbox.models.tokenizers import MTLTokenizer from chatterbox.models.voice_encoder import VoiceEncoder from chatterbox.mtl_tts import ChatterboxMultilingualTTS except ModuleNotFoundError: # Compatibility with the public Chatterbox fine-tuning kit layout. from src.chatterbox_.models.s3gen import S3Gen from src.chatterbox_.models.t3 import T3 from src.chatterbox_.models.t3.modules.t3_config import T3Config from src.chatterbox_.models.tokenizers import MTLTokenizer from src.chatterbox_.models.voice_encoder import VoiceEncoder from src.chatterbox_.mtl_tts import ChatterboxMultilingualTTS def load_model(model_dir: Path, device: str) -> ChatterboxMultilingualTTS: voice_encoder = VoiceEncoder() voice_encoder.load_state_dict(load_file(str(model_dir / "ve.safetensors"))) voice_encoder.to(device).eval() t3 = T3(T3Config.multilingual()) t3.load_state_dict(load_file(str(model_dir / "t3_fa.safetensors"))) t3.tfmr.set_attn_implementation("eager") t3.to(device).eval() s3gen = S3Gen() missing, unexpected = s3gen.load_state_dict( load_file(str(model_dir / "s3gen.safetensors")), strict=False ) if unexpected or missing not in (["tokenizer.window"], []): raise RuntimeError( f"decoder state mismatch: missing={missing}, unexpected={unexpected}" ) s3gen.to(device).eval() tokenizer = MTLTokenizer( str(model_dir / "grapheme_mtl_merged_expanded_v1.json") ) return ChatterboxMultilingualTTS( t3=t3, s3gen=s3gen, ve=voice_encoder, tokenizer=tokenizer, device=device, ) def main() -> int: parser = argparse.ArgumentParser() parser.add_argument("--model-dir", type=Path, required=True) parser.add_argument("--text", required=True) parser.add_argument("--reference", type=Path, required=True) parser.add_argument("--output", type=Path, required=True) parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") parser.add_argument("--seed", type=int, default=4200) parser.add_argument("--exaggeration", type=float, default=0.5) parser.add_argument("--cfg-weight", type=float, default=0.5) args = parser.parse_args() torch.manual_seed(args.seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(args.seed) model = load_model(args.model_dir, args.device) wav = model.generate( text=args.text, language_id=None, audio_prompt_path=str(args.reference), exaggeration=args.exaggeration, cfg_weight=args.cfg_weight, ) args.output.parent.mkdir(parents=True, exist_ok=True) sf.write(args.output, wav.squeeze().cpu().numpy(), model.sr) return 0 if __name__ == "__main__": raise SystemExit(main())