Kazeia-engine/dist/distill/prep_data.py

100 lines
4.1 KiB
Python

# Prep conditions de distillation depuis un corpus audio FR (pas de texte requis).
# Par utterance: speech tokens 25Hz (s3tok ONNX), mel 50Hz (matcha 1920/480@24k), x-vector campplus.
# Sortie: shards .pt (listes de dicts) dans --out_dir.
# Usage: python prep_data.py --corpus_dir data/cml --model_dir models/cosyvoice3-0.5b --out_dir shards [--max_utts N]
import argparse, glob, os, sys
import numpy as np
import torch
import torchaudio
import torchaudio.compliance.kaldi as kaldi
import onnxruntime as ort
import whisper
from matcha.utils.audio import mel_spectrogram
MIN_S, MAX_S = 6.0, 28.0
def load_resample(path, sr):
# soundfile (torchaudio>=2.9 exige torchcodec pour load)
import soundfile as sf
data, fs = sf.read(path, dtype="float32", always_2d=True)
wav = torch.from_numpy(data.T)
if wav.shape[0] > 1:
wav = wav.mean(0, keepdim=True)
if fs != sr:
wav = torchaudio.functional.resample(wav, fs, sr)
return wav
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--corpus_dir", required=True)
ap.add_argument("--model_dir", required=True)
ap.add_argument("--out_dir", required=True)
ap.add_argument("--max_utts", type=int, default=0)
ap.add_argument("--shard_size", type=int, default=500)
ap.add_argument("--device", default="cuda")
args = ap.parse_args()
os.makedirs(args.out_dir, exist_ok=True)
prov = ["CUDAExecutionProvider", "CPUExecutionProvider"] if args.device == "cuda" else ["CPUExecutionProvider"]
s3tok = ort.InferenceSession(f"{args.model_dir}/speech_tokenizer_v3.onnx", providers=prov)
camp = ort.InferenceSession(f"{args.model_dir}/campplus.onnx", providers=["CPUExecutionProvider"])
files = []
for ext in ("wav", "flac", "mp3", "ogg"):
files += glob.glob(f"{args.corpus_dir}/**/*.{ext}", recursive=True)
files.sort()
if args.max_utts:
files = files[: args.max_utts]
print(f"{len(files)} fichiers candidats", flush=True)
shard, n_shard, n_ok = [], 0, 0
for i, f in enumerate(files):
try:
w16 = load_resample(f, 16000)
dur = w16.shape[1] / 16000
if not (MIN_S <= dur <= MAX_S):
continue
w24 = load_resample(f, 24000)
# tokens 25Hz
feat = whisper.log_mel_spectrogram(w16, n_mels=128)
toks = s3tok.run(None, {s3tok.get_inputs()[0].name: feat.numpy(),
s3tok.get_inputs()[1].name: np.array([feat.shape[2]], dtype=np.int32)})[0].flatten()
# mel 50Hz
mel = mel_spectrogram(w24, n_fft=1920, num_mels=80, sampling_rate=24000,
hop_size=480, win_size=1920, fmin=0, fmax=None, center=False)
mel = mel.squeeze(0).transpose(0, 1) # [T,80]
# alignement mel = 2 * tokens
tl = min(mel.shape[0] // 2, len(toks))
if tl < int(MIN_S * 25):
continue
mel, toks = mel[: 2 * tl], toks[:tl]
# x-vector
fb = kaldi.fbank(w16, num_mel_bins=80, dither=0, sample_frequency=16000)
fb = fb - fb.mean(dim=0, keepdim=True)
xv = camp.run(None, {camp.get_inputs()[0].name: fb.unsqueeze(0).numpy()})[0].flatten()
shard.append({"path": f,
"tokens": torch.tensor(toks.astype(np.int32), dtype=torch.int32),
"mel": mel.to(torch.float16),
"xvec": torch.tensor(xv, dtype=torch.float16)})
n_ok += 1
if len(shard) >= args.shard_size:
torch.save(shard, f"{args.out_dir}/shard_{n_shard:05d}.pt")
n_shard += 1
shard = []
if n_ok % 200 == 0:
print(f"[{i}/{len(files)}] {n_ok} ok, {n_shard} shards", flush=True)
except Exception as e:
print(f"skip {f}: {e}", file=sys.stderr, flush=True)
if shard:
torch.save(shard, f"{args.out_dir}/shard_{n_shard:05d}.pt")
n_shard += 1
print(f"FINI: {n_ok} utts -> {n_shard} shards dans {args.out_dir}", flush=True)
if __name__ == "__main__":
main()