malagasy-tts / reference_encoding.py
mimba's picture
Add application file
ab7b3be
Raw
History Blame Contribute Delete
8.03 kB
"""
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)