File size: 4,606 Bytes
e352304
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5da3b02
e352304
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dee8fe1
 
 
 
e352304
 
 
 
5da3b02
 
 
 
 
 
 
 
 
 
 
 
 
e352304
dee8fe1
e352304
 
 
 
 
 
 
 
 
 
 
 
 
 
dee8fe1
 
 
 
e352304
 
 
dee8fe1
 
e352304
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""RegaLabs-TTS: Direct Hugging Face Sorani Text-to-Speech Inference Script."""

import argparse
import sys
import os
from pathlib import Path
import torch
import soundfile as sf

def parse_args():
    parser = argparse.ArgumentParser(description="RegaLabs-TTS Sorani Speech Synthesis")
    parser.add_argument("--text", required=True, help="Sorani text to synthesize")
    parser.add_argument("--prompt-wav", required=True, help="Path to reference audio WAV")
    parser.add_argument("--prompt-text", required=True, help="Transcript of reference audio prompt")
    parser.add_argument("--flow-checkpoint", default=str(Path(__file__).parent / "cosyvoice3_sorani_flow_best_step2300.pt"), help="Path to flow checkpoint")
    parser.add_argument("--adapter", default=str(Path(__file__).parent / "cosyvoice3_sorani_lora_refined_best.pt"), help="Path to the Sorani LLM LoRA adapter (.pt)")
    parser.add_argument("--base-model", default="FunAudioLLM/Fun-CosyVoice3-0.5B-2512", help="Hugging Face base model ID")
    parser.add_argument("--out", default="output_sorani.wav", help="Output WAV path")
    parser.add_argument("--cosyvoice-repo", default=os.environ.get("COSYVOICE_REPO", "./CosyVoice"), help="Path to cloned CosyVoice repo")
    return parser.parse_args()

def main():
    args = parse_args()
    
    repo = Path(args.cosyvoice_repo).resolve()
    if not repo.exists():
        print(f"Error: CosyVoice engine directory not found at {repo}.", file=sys.stderr)
        print("Please clone CosyVoice: git clone --recursive https://github.com/FunAudioLLM/CosyVoice.git", file=sys.stderr)
        sys.exit(1)
        
    sys.path.insert(0, str(repo))
    matcha_path = repo / "third_party" / "Matcha-TTS"
    if matcha_path.exists():
        sys.path.insert(0, str(matcha_path))
        
    sys.path.insert(0, str(Path(__file__).parent))
    from sorani.frontend import normalize_sorani_text
    from sorani.censor import verify_checkpoint

    from cosyvoice.cli.cosyvoice import CosyVoice3
    
    print(f"Loading base model: {args.base_model}...")
    cosyvoice = CosyVoice3(args.base_model, fp16=torch.cuda.is_available())

    adapter = Path(args.adapter)
    if adapter.exists():
        from cosyvoice.utils.lora import inject_lora, load_lora_state_dict
        llm = cosyvoice.model.llm
        target_count = inject_lora(llm, rank=16, alpha=32.0, dropout=0.05)
        load_lora_state_dict(
            llm, torch.load(adapter, map_location="cpu", weights_only=False)
        )
        llm.to(cosyvoice.model.device).eval()
        print(f"Loaded RegaLabs-TTS Sorani LLM adapter ({target_count} projections).")
    elif str(adapter) != str(Path(__file__).parent / "cosyvoice3_sorani_lora_refined_best.pt"):
        print(f"Warning: adapter not found at {adapter}; continuing without LLM adapter.", file=sys.stderr)
    
    verify_checkpoint(args.flow_checkpoint)
    print(f"Loading RegaLabs-TTS Sorani flow checkpoint: {args.flow_checkpoint}...")
    flow_state = torch.load(args.flow_checkpoint, map_location="cpu", weights_only=False)
    if isinstance(flow_state, dict):
        for key in ("model", "state_dict", "flow"):
            nested = flow_state.get(key)
            if isinstance(nested, dict):
                flow_state = nested
                break
        flow_state = {k: v for k, v in flow_state.items() if isinstance(k, str) and isinstance(v, torch.Tensor)}
        
    cosyvoice.model.flow.load_state_dict(flow_state, strict=False)
    cosyvoice.model.flow.to(cosyvoice.model.device).eval()
    
    print(f"Synthesizing Sorani text: '{args.text}'...")
    clean_text = normalize_sorani_text(args.text)
    clean_prompt_text = normalize_sorani_text(args.prompt_text)
    if "<|endofprompt|>" not in clean_prompt_text:
        clean_prompt_text = f"You are a helpful assistant.<|endofprompt|> {clean_prompt_text}"
    pieces = []
    with torch.inference_mode():
        for output in cosyvoice.inference_zero_shot(
            clean_text,
            clean_prompt_text,
            args.prompt_wav,
            stream=False,
            text_frontend=False,
        ):
            pieces.append(output["tts_speech"].detach().cpu())
            
    if not pieces:
        raise RuntimeError("No audio output returned.")
        
    speech = torch.cat(pieces, dim=1).squeeze(0).numpy()
    out_path = Path(args.out)
    out_path.parent.mkdir(parents=True, exist_ok=True)
    sf.write(out_path, speech, cosyvoice.sample_rate, subtype="PCM_16")
    print(f"✅ Audio generated successfully: {out_path}")

if __name__ == "__main__":
    main()