File size: 8,522 Bytes
df972d3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2ca8adc
df972d3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2ca8adc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
df972d3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
import os
import sys
import argparse
import torch
import soundfile as sf

CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
TTS_ROOT = os.path.abspath(os.path.join(CURRENT_DIR, ".."))
if TTS_ROOT not in sys.path:
    sys.path.insert(0, TTS_ROOT)

# HuggingFace Cache ve Proxy ayarları (Ağ bloklanması / takılmasını önlemek için)
if "HF_HOME" not in os.environ:
    scratch_hf = "/gpfs/scratch/ehpc540/users/huseyin/.cache/huggingface"
    if os.path.exists(scratch_hf):
        os.environ["HF_HOME"] = scratch_hf

if "HTTP_PROXY" not in os.environ and "http_proxy" not in os.environ:
    os.environ["HTTP_PROXY"] = "http://127.0.0.1:9595"
    os.environ["HTTPS_PROXY"] = "http://127.0.0.1:9595"
    os.environ["http_proxy"] = "http://127.0.0.1:9595"
    os.environ["https_proxy"] = "http://127.0.0.1:9595"

try:
    from model.model import TTSModel, TTSConfig
    from dataset.tokenizer import TextTokenizer, AudioCodecTokenizer, TTSProcessor
except (ImportError, ModuleNotFoundError):
    from model import TTSModel, TTSConfig
    from tokenizer import TextTokenizer, AudioCodecTokenizer, TTSProcessor


import numpy as np

def trim_silence_and_noise(wav: np.ndarray, sr: int = 24000, frame_ms: int = 50, pad_ms: int = 250) -> np.ndarray:
    """
    Konuşma bittikten sonra autoregressive modelin ürettiği boşluk veya cızırtıyı (statik gürültü)
    kısa süreli RMS enerji analiziyle tespit edip yumuşak bir fade-out ile temizler.
    """
    frame_len = int(sr * (frame_ms / 1000.0))
    num_frames = len(wav) // frame_len
    if num_frames == 0:
        return wav

    frames = [wav[i * frame_len : (i + 1) * frame_len] for i in range(num_frames)]
    energies = [np.sqrt(np.mean(f**2)) for f in frames]
    peak_energy = max(energies) if len(energies) > 0 else 0.0

    if peak_energy < 0.01:
        return wav

    # Aktif konuşma eşiği: tepe enerjinin %20'si (en az 0.032)
    thresh = max(0.032, peak_energy * 0.20)

    last_active_frame = 0
    for i in range(num_frames - 1, -1, -1):
        if energies[i] > thresh:
            last_active_frame = i
            break

    # Konuşmanın doğal sönümlenmesi için 250ms tampon (pad) bırak
    pad_frames = int(pad_ms / frame_ms)
    cut_frame = min(num_frames, last_active_frame + 1 + pad_frames)
    cut_sample = min(len(wav), cut_frame * frame_len)

    trimmed = wav[:cut_sample].copy()
    # Tıklama/pop sesini önlemek için son 50ms'ye yumuşak fade-out
    fade_len = min(int(sr * 0.05), len(trimmed))
    if fade_len > 0:
        fade_curve = np.linspace(1.0, 0.0, fade_len)
        trimmed[-fade_len:] *= fade_curve

    cut_sec = (len(wav) - len(trimmed)) / sr
    if cut_sec > 0.3:
        print(f"[Inference] ✂️ Akıllı Kırpıcı: Orijinal {len(wav)/sr:.2f}s -> Temizlenmiş {len(trimmed)/sr:.2f}s ({cut_sec:.2f} saniye cızırtı/statik kesildi)", flush=True)

    return trimmed


def main():
    parser = argparse.ArgumentParser(description="TTS Inference and Audio Synthesis Pipeline")
    parser.add_argument("--text", type=str, required=True, help="Seslendirilecek metin")
    parser.add_argument("--instruction", type=str, default=None, help="Ses tasarımı talimatı (Voice Design - Örn: 'Kalın ve tok bir ses')")
    parser.add_argument("--ref_audio", type=str, default=None, help="Klonlanacak referans ses dosyası (Voice Clone - .wav)")
    parser.add_argument("--ref_max_sec", type=float, default=None, help="Referans sesin maksimum süresi (Varsayılan: None - sesin tamamı kullanılır)")
    parser.add_argument("--checkpoint", type=str, default="model.safetensors", help="Model ağırlık dosyası (.safetensors veya .pt)")
    parser.add_argument("--use_qwen_backbone", action="store_true", help="Qwen2.5-0.5B pretrained ağırlıklarını omurgaya yükler")
    parser.add_argument("--output", type=str, default="output.wav", help="Kaydedilecek ses dosyası yolu")
    parser.add_argument("--max_tokens", type=int, default=120, help="Maksimum üretilecek ses karesi (12.5 Hz'de 120 kare ~ 9.6 saniye)")
    parser.add_argument("--temperature", type=float, default=0.6, help="Sampling sıcaklığı (0.6 doğal tonlama ve melodi sağlar)")
    parser.add_argument("--depth_temperature", type=float, default=0.6, help="Akustik katman sıcaklığı (Yumuşak sampling ile robotikliği kırar)")
    parser.add_argument("--top_k", type=int, default=30, help="Top-k sampling parametresi")
    parser.add_argument("--no_trim", action="store_true", help="Cümle sonundaki sessizlik/cızırtı kırpmasını devre dışı bırakır")
    parser.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu")
    args = parser.parse_args()

    device = torch.device(args.device)
    print(f"[Inference] Cihaz: {device}")
    print(f"[Inference] Metin: '{args.text}'")
    if args.instruction:
        print(f"[Inference] Voice Design Talimatı: '{args.instruction}'")
    if args.ref_audio:
        print(f"[Inference] Voice Clone Referansı: '{args.ref_audio}'")

    # 1. Model ve Tokenizer Yükleme
    if args.use_qwen_backbone:
        print("[Inference] Qwen2.5-0.5B omurgası ile model konfigüre ediliyor...")
        config = TTSConfig.from_qwen("Qwen/Qwen2.5-0.5B")
        model = TTSModel(config).to(device)
        if not args.checkpoint:
            model.load_qwen_backbone("Qwen/Qwen2.5-0.5B")
    else:
        config = TTSConfig()
        model = TTSModel(config).to(device)

    ckpt_path = args.checkpoint
    if not ckpt_path or not os.path.exists(ckpt_path):
        if os.path.exists("model.safetensors"):
            ckpt_path = "model.safetensors"
        elif os.path.exists("model.pt"):
            ckpt_path = "model.pt"

    if ckpt_path and os.path.exists(ckpt_path):
        print(f"[Inference] Checkpoint yükleniyor: {ckpt_path}", flush=True)
        if ckpt_path.endswith(".safetensors"):
            from safetensors.torch import load_file
            state_dict = load_file(ckpt_path)
            state_dict = {k: v.to(device) for k, v in state_dict.items()}
        else:
            ckpt = torch.load(ckpt_path, map_location=device)
            state_dict = ckpt.get("model", ckpt) if isinstance(ckpt, dict) else ckpt
        if any(k.startswith("module.") for k in state_dict.keys()):
            state_dict = {k.replace("module.", ""): v for k, v in state_dict.items()}
        model.load_state_dict(state_dict)
        print("[Inference] Model ağırlıkları başarıyla yüklendi!", flush=True)
    elif not args.use_qwen_backbone:
        print("[Inference] Bilgi: Checkpoint verilmedi, model başlangıç ağırlıklarıyla çalışacak.", flush=True)

    print("[Inference] 1/4 Metin tokenizer hazırlanıyor...", flush=True)
    text_tokenizer = TextTokenizer()
    print("[Inference] 2/4 Kyutai Mimi ses çözücü (Audio Codec) yükleniyor...", flush=True)
    audio_tokenizer = AudioCodecTokenizer(device=args.device)
    processor = TTSProcessor(text_tokenizer, audio_tokenizer)

    # 2. Girdileri Hazırla
    print("[Inference] 3/4 Girdiler hazırlanıyor...", flush=True)
    text_ids, ref_codes = processor.prepare_inference_inputs(
        text=args.text,
        instruction=args.instruction,
        ref_audio_path=args.ref_audio,
        max_ref_sec=args.ref_max_sec,
        device=args.device,
    )

    # 3. Autoregressive Üretim (CB0 + CB1..7)
    print(f"[Inference] 4/4 Ses tokenları üretiliyor (Maksimum {args.max_tokens} kare)...", flush=True)
    generated_codes = model.generate(
        text_ids=text_ids,
        ref_audio_codes=ref_codes,
        max_new_tokens=args.max_tokens,
        temperature=args.temperature,
        depth_temperature=args.depth_temperature,
        top_k=args.top_k,
    )
    print(f"[Inference] Üretilen ses token matrisi boyutu: {generated_codes.shape}", flush=True)

    # 4. Codec ile Sese Dönüştürme
    print("[Inference] Ses dalgasına (waveform / 24 kHz) dönüştürülüyor...", flush=True)
    wav = audio_tokenizer.decode(generated_codes)
    if isinstance(wav, torch.Tensor):
        wav = wav.squeeze().cpu().numpy()

    # 5. Otomatik Cızırtı ve Sessizlik Kırpma (VAD / Energy Trimmer)
    if not args.no_trim:
        wav = trim_silence_and_noise(wav, sr=audio_tokenizer.sample_rate)

    # 6. Dosyaya Yazma
    sf.write(args.output, wav, samplerate=audio_tokenizer.sample_rate)
    print(f"[Inference] 🎉 Başarılı! Temiz ve optimize edilmiş ses kaydedildi -> {args.output}", flush=True)


if __name__ == "__main__":
    main()