Download inference.py from TurkishCodeMan/NeuroVoice-0.5B: direct link, hf CLI and curl.
- Browser
- Download file 8.52 kB
-
https://huggingface.co/TurkishCodeMan/NeuroVoice-0.5B/resolve/main/inference.py
- Command line
-
hf download hf://TurkishCodeMan/NeuroVoice-0.5B/inference.py
-
curl -L -o inference.py https://huggingface.co/TurkishCodeMan/NeuroVoice-0.5B/resolve/main/inference.py
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() | |