Spaces:
Running
Running
| """ | |
| 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) | |