# 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()