NeuroVoice-0.5B / inference.py
TurkishCodeMan's picture
Upload inference.py with huggingface_hub
2ca8adc verified
Raw History Blame Contribute Delete
8.52 kB
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()