100 lines
4.1 KiB
Python
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()
|