""" Encodage d'une reference audio ARBITRAIRE (clonage zero-shot) -- style_ttl/style_dp calcules a la volee via les graphes ONNX dedies (codec_encoder, style_encoder, duration_style_encoder), exportes avec le modele (voir notebook d'entrainement section 19 -- exports/export_onnx.py, --slim). Reproduit fidelement export_new_voice.py (script officiel de production, verifie ligne par ligne) : nettoyage audio, reechantillonnage haute qualite (memes parametres exacts), troncature du mel a un multiple exact de chunk_compress_factor. NOTE IMPORTANTE (voir notes projet) : le clonage FONCTIONNE techniquement, mais le modele actuel (finetune sur seulement 4 locuteurs) a tendance a faire converger une reference arbitraire vers l'un de ces 4 locuteurs d'entrainement, plutot que de vraiment cloner le timbre fourni -- limitation connue, a corriger par un futur reentrainement avec plus de diversite de locuteurs. Cette implementation est deliberement complete des maintenant pour ne pas avoir a retoucher l'architecture du worker quand ce probleme sera corrige. """ import numpy as np import torch import torchaudio def rechantillonner_haute_qualite(wav, sr_in, sr_out, device=None): """Copie EXACTE de ensure_sr() dans export_new_voice.py (production).""" if device is None: device = wav.device if wav.dim() == 1: wav = wav.unsqueeze(0) if sr_in != sr_out: wav = torchaudio.functional.resample( wav, sr_in, sr_out, lowpass_filter_width=64, rolloff=0.9475937167399596, resampling_method="sinc_interp_kaiser", beta=14.769656459379492, ) return wav.to(device) def nettoyer_audio_reference(wav, sr, strength=0.75, noise_floor_quantile=0.20): """Copie EXACTE de clean_reference_audio() dans export_new_voice.py (production).""" if wav.numel() == 0: return wav strength = float(max(0.0, min(strength, 1.0))) if strength <= 0.0: return wav wav = torch.nan_to_num(wav.to(torch.float32)) if wav.dim() == 1: wav = wav.unsqueeze(0) original = wav.clone() wav = wav - wav.mean(dim=-1, keepdim=True) try: wav = torchaudio.functional.highpass_biquad(wav, sr, cutoff_freq=70.0) wav = torchaudio.functional.lowpass_biquad(wav, sr, cutoff_freq=min(8000.0, sr * 0.45)) except Exception: pass length = wav.shape[-1] n_fft = min(2048, max(256, 2 ** int(np.floor(np.log2(max(length, 256)))))) hop = max(1, n_fft // 4) if length < n_fft: return wav window = torch.hann_window(n_fft, device=wav.device) spec = torch.stft(wav.squeeze(0), n_fft=n_fft, hop_length=hop, win_length=n_fft, window=window, return_complex=True) mag = spec.abs() frame_energy = mag.mean(dim=0) q = float(max(0.05, min(noise_floor_quantile, 0.5))) threshold_energy = torch.quantile(frame_energy, q) noise_frames = frame_energy <= threshold_energy noise_profile = mag[:, noise_frames].median(dim=1).values if bool(noise_frames.any()) else mag.median(dim=1).values threshold = noise_profile[:, None] * (1.2 + 1.4 * strength) attenuation_floor = 1.0 - (0.85 * strength) gain = ((mag - threshold) / (mag + 1e-8)).clamp(min=attenuation_floor, max=1.0) cleaned = torch.istft(spec * gain, n_fft=n_fft, hop_length=hop, win_length=n_fft, window=window, length=length).unsqueeze(0) orig_rms = torch.sqrt(torch.mean(original.square())).clamp_min(1e-6) clean_rms = torch.sqrt(torch.mean(cleaned.square())).clamp_min(1e-6) cleaned = cleaned * torch.clamp(orig_rms / clean_rms, max=2.0) peak = cleaned.abs().max().clamp_min(1e-6) if peak > 0.98: cleaned = cleaned / peak * 0.98 return cleaned class LinearMelSpectrogram(torch.nn.Module): """Copie EXACTE de LinearMelSpectrogram (models/utils.py, production) -- concatene log-spectrogramme lineaire ET log-mel (pas juste un mel classique).""" def __init__(self, sample_rate=44100, n_fft=2048, win_length=2048, hop_length=512, n_mels=228, f_min=0, f_max=None): super().__init__() self.spectrogram = torchaudio.transforms.Spectrogram( n_fft=n_fft, win_length=win_length, hop_length=hop_length, center=True, power=1.0, ) self.mel_scale = torchaudio.transforms.MelScale( n_mels=n_mels, sample_rate=sample_rate, n_stft=n_fft // 2 + 1, f_min=f_min, f_max=f_max, ) def forward(self, audio): spec = self.spectrogram(audio) mel = self.mel_scale(spec) log_spec = torch.log(torch.clamp(spec, min=1e-5)) log_mel = torch.log(torch.clamp(mel, min=1e-5)) return torch.cat([log_spec, log_mel], dim=1) def build_reference_only(z_ref_input, valid_len, max_frames=None): """Copie EXACTE de build_reference_only() (training/t2l/sampling.py) -- reimplementee en numpy pur (pas besoin du depot d'entrainement complet).""" B, C, T_ref = z_ref_input.shape if max_frames is not None and T_ref > max_frames: z_ref_input = z_ref_input[:, :, :max_frames] T_ref = max_frames arange_T = np.arange(T_ref).reshape(1, -1) valid_len_c = np.clip(valid_len, 0, T_ref).reshape(-1, 1) ref_mask = (arange_T < valid_len_c).reshape(B, 1, T_ref).astype(np.float32) z_ref_left = z_ref_input * ref_mask return z_ref_left, ref_mask class EncodeurReference: """Charge les 3 graphes ONNX zero-shot (codec_encoder, style_encoder, duration_style_encoder) et expose encoder(chemin_wav) -> (style_ttl, style_dp), entierement fidele au pipeline officiel de production.""" def __init__(self, onnx_dir, sample_rate=44100, n_fft=2048, hop_length=512, n_mels=228, chunk_compress_factor=6, providers=("CPUExecutionProvider",)): import os import onnxruntime as ort self.sample_rate = sample_rate self.chunk_compress_factor = chunk_compress_factor self.mel_spec = LinearMelSpectrogram( sample_rate=sample_rate, n_fft=n_fft, hop_length=hop_length, n_mels=n_mels, ) self.codec_encoder = ort.InferenceSession(os.path.join(onnx_dir, "codec_encoder.onnx"), providers=list(providers)) self.style_encoder = ort.InferenceSession(os.path.join(onnx_dir, "style_encoder.onnx"), providers=list(providers)) self.duration_style_encoder = ort.InferenceSession( os.path.join(onnx_dir, "duration_style_encoder.onnx"), providers=list(providers)) def encoder(self, chemin_wav, nettoyer=True): import soundfile as sf wav_np, sr_ref = sf.read(chemin_wav) if wav_np.ndim > 1: wav_np = wav_np.mean(axis=1) # downmix mono, meme principe que read_wav_mono() wav_torch = torch.tensor(wav_np, dtype=torch.float32).unsqueeze(0) if sr_ref != self.sample_rate: wav_torch = rechantillonner_haute_qualite(wav_torch, sr_ref, self.sample_rate) if nettoyer: wav_torch = nettoyer_audio_reference(wav_torch, self.sample_rate) with torch.no_grad(): mel = self.mel_spec(wav_torch) tm = mel.shape[-1] aligned = (tm // self.chunk_compress_factor) * self.chunk_compress_factor if aligned < tm: mel = mel[..., :aligned] mel_np = mel.numpy().astype(np.float32) z_ref, *_ = self.codec_encoder.run(None, {"mel": mel_np}) z_ref = np.asarray(z_ref, dtype=np.float32) valid_len = np.array([z_ref.shape[2]], dtype=np.int64) z_tr, ref_mask = build_reference_only(z_ref, valid_len, max_frames=None) style_ttl, *_ = self.style_encoder.run(None, {"z_ref": z_tr, "ref_mask": ref_mask}) style_dp, *_ = self.duration_style_encoder.run(None, {"z_ref": z_tr, "ref_mask": ref_mask}) return np.asarray(style_ttl, dtype=np.float32), np.asarray(style_dp, dtype=np.float32)