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)