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