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