Spaces:
Running
Running
File size: 8,030 Bytes
ab7b3be | 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 | """
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)
|