| |
| """LEVER C — fine-tune du JOINT w2v-BERT (joint_cont) avec AUGMENTATION anti-sur-apprentissage OOD. |
| Adapte train_mms_adapter_aug.py au modele Wav2Vec2Bert : |
| - REUTILISE le vocab/processor du checkpoint init (joint_cont) => tete CTC conservee. |
| - Feature extractor SeamlessM4T (input_features, 80 mel), 25 fps. |
| - Full fine-tune BAS LR (pas d'adaptateurs MMS) + SpecAugment agressif (config modele). |
| - AUGMENTATION audio (audiomentations) sur TRAIN uniquement : bruit, time-stretch, pitch, gain. |
| - eval = dev-difficile (locuteurs DISJOINTS), NON augmente. 100% conforme (aucune donnee externe). |
| """ |
| import argparse, json, os, random |
| from dataclasses import dataclass |
|
|
| import jiwer, numpy as np, soundfile as sf, torch |
| from datasets import Dataset |
| from transformers import Trainer, TrainingArguments, AutoModelForCTC, AutoProcessor |
| from transformers.trainer_pt_utils import LengthGroupedSampler |
| from audiomentations import Compose, AddGaussianNoise, TimeStretch, PitchShift, Gain |
|
|
| SR = 16000 |
| FRAMES_PER_SEC = 25 |
|
|
|
|
| class LGTrainer(Trainer): |
| def __init__(self, *args, eval_collator=None, **kwargs): |
| super().__init__(*args, **kwargs) |
| self._eval_collator = eval_collator |
|
|
| def _get_train_sampler(self, train_dataset=None): |
| ds = train_dataset if train_dataset is not None else self.train_dataset |
| return LengthGroupedSampler( |
| self.args.per_device_train_batch_size * self.args.gradient_accumulation_steps, |
| dataset=ds, lengths=list(ds["length"])) |
|
|
| def get_eval_dataloader(self, eval_dataset=None): |
| if self._eval_collator is None: |
| return super().get_eval_dataloader(eval_dataset) |
| saved = self.data_collator |
| self.data_collator = self._eval_collator |
| try: |
| return super().get_eval_dataloader(eval_dataset) |
| finally: |
| self.data_collator = saved |
|
|
|
|
| def read_jsonl(p): |
| rows = [] |
| for l in open(p, encoding="utf-8"): |
| r = json.loads(l) |
| r["text"] = " ".join(r["text"].replace("|", " ").split()) |
| rows.append(r) |
| return rows |
|
|
|
|
| def parse(): |
| p = argparse.ArgumentParser() |
| p.add_argument("--out", required=True) |
| p.add_argument("--train", nargs="+", required=True) |
| p.add_argument("--eval", required=True) |
| p.add_argument("--init", required=True) |
| p.add_argument("--lr", type=float, default=2e-5) |
| p.add_argument("--epochs", type=float, default=2) |
| p.add_argument("--bs", type=int, default=8) |
| p.add_argument("--grad_accum", type=int, default=4) |
| p.add_argument("--warmup", type=int, default=200) |
| p.add_argument("--eval_steps", type=int, default=150) |
| p.add_argument("--eval_subset", type=int, default=2000) |
| p.add_argument("--min_dur", type=float, default=1.0) |
| p.add_argument("--max_dur", type=float, default=25.0) |
| p.add_argument("--max_wps", type=float, default=4.5) |
| p.add_argument("--num_workers", type=int, default=8) |
| p.add_argument("--save_total_limit", type=int, default=2) |
| p.add_argument("--augment", action="store_true") |
| p.add_argument("--mask_time_prob", type=float, default=0.075) |
| p.add_argument("--mask_feature_prob", type=float, default=0.05) |
| p.add_argument("--freeze_frontend", action="store_true") |
| p.add_argument("--smoke", action="store_true") |
| p.add_argument("--seed", type=int, default=42) |
| return p.parse_args() |
|
|
|
|
| def main(): |
| a = parse() |
| os.makedirs(a.out, exist_ok=True) |
| random.seed(a.seed); np.random.seed(a.seed); torch.manual_seed(a.seed) |
|
|
| |
| processor = AutoProcessor.from_pretrained(a.init) |
| tok = processor.tokenizer |
| fe = processor.feature_extractor |
| processor.save_pretrained(a.out) |
|
|
| train_rows = [] |
| for m in a.train: |
| rows = read_jsonl(m) |
| kept = [r for r in rows if r["text"] and a.min_dur <= r["duration"] <= a.max_dur |
| and len(r["text"].split()) / max(r["duration"], 0.1) <= a.max_wps] |
| print(f"{m}: {len(kept)}/{len(rows)} gardes ({sum(r['duration'] for r in kept)/3600:.1f}h)", flush=True) |
| train_rows += kept |
| eval_rows = [r for r in read_jsonl(a.eval) if r["text"]] |
| rng = random.Random(a.seed) |
| if a.eval_subset and len(eval_rows) > a.eval_subset: |
| eval_rows = rng.sample(eval_rows, a.eval_subset) |
| if a.smoke: |
| train_rows, eval_rows = train_rows[:64], eval_rows[:24] |
| a.epochs, a.eval_steps, a.warmup, a.bs, a.grad_accum = 1, 5, 2, 4, 1 |
| print(f"TRAIN {len(train_rows)} | EVAL {len(eval_rows)} | augment={a.augment}", flush=True) |
|
|
| def feasible(r): |
| ids = tok(r["text"]).input_ids |
| need = len(ids) + sum(x == y for x, y in zip(ids, ids[1:])) |
| return need <= int(r["duration"] * FRAMES_PER_SEC) - 2 |
| before = len(train_rows) |
| train_rows = [r for r in train_rows if feasible(r)] |
| print(f"feasible CTC: {len(train_rows)}/{before}", flush=True) |
|
|
| def to_ds(rows): |
| return Dataset.from_list([{"audio": r["audio"], "text": r["text"], |
| "length": int(r["duration"] * 100)} for r in rows]) |
| train_ds, eval_ds = to_ds(train_rows), to_ds(eval_rows) |
|
|
| augpipe = None |
| if a.augment: |
| augpipe = Compose([ |
| AddGaussianNoise(min_amplitude=0.001, max_amplitude=0.015, p=0.5), |
| TimeStretch(min_rate=0.9, max_rate=1.1, leave_length_unchanged=False, p=0.4), |
| PitchShift(min_semitones=-2.0, max_semitones=2.0, p=0.25), |
| Gain(min_gain_db=-6.0, max_gain_db=6.0, p=0.4), |
| ]) |
|
|
| @dataclass |
| class Collator: |
| augment: object = None |
| def __call__(self, feats): |
| audio = [] |
| for f in feats: |
| w = sf.read(f["audio"], dtype="float32")[0] |
| if self.augment is not None: |
| try: |
| w = self.augment(samples=w, sample_rate=SR) |
| except Exception: |
| pass |
| audio.append(w) |
| batch = fe(audio, sampling_rate=SR, return_tensors="pt", padding=True) |
| enc = tok([f["text"] for f in feats], return_tensors="pt", padding=True) |
| batch["labels"] = enc["input_ids"].masked_fill(enc["attention_mask"].ne(1), -100) |
| return batch |
|
|
| train_collator = Collator(augment=augpipe) |
| eval_collator = Collator(augment=None) |
|
|
| model = AutoModelForCTC.from_pretrained( |
| a.init, ctc_loss_reduction="mean", ctc_zero_infinity=True, |
| apply_spec_augment=True, |
| mask_time_prob=a.mask_time_prob, mask_time_length=10, |
| mask_feature_prob=a.mask_feature_prob, mask_feature_length=10) |
| if a.freeze_frontend and hasattr(model, "wav2vec2_bert"): |
| |
| for p in model.wav2vec2_bert.feature_projection.parameters(): |
| p.requires_grad = False |
| n_train = sum(p.numel() for p in model.parameters() if p.requires_grad) |
| print(f"params entrainables: {n_train/1e6:.1f}M", flush=True) |
|
|
| def preprocess_logits(logits, labels): |
| return torch.argmax(logits, dim=-1) |
|
|
| eval_refs = [r["text"] for r in eval_rows] |
|
|
| def metrics(pred): |
| ids = np.where(pred.predictions == -100, tok.pad_token_id, pred.predictions) |
| hyps = tok.batch_decode(ids) |
| for _r, _h in list(zip(eval_refs, hyps))[:3]: |
| print(" [dbg] REF:", _r[:55], "|| HYP:", repr(_h[:55]), flush=True) |
| pairs = [(r, h) for r, h in zip(eval_refs, hyps) if r.strip()] |
| refs = [r for r, _ in pairs]; hs = [h for _, h in pairs] |
| wer = jiwer.wer(refs, hs); cer = jiwer.cer(refs, hs) |
| return {"wer": wer, "cer": cer, "combine": 0.5 * wer + 0.5 * cer} |
|
|
| targs = TrainingArguments( |
| output_dir=a.out, per_device_train_batch_size=a.bs, per_device_eval_batch_size=a.bs, |
| gradient_accumulation_steps=a.grad_accum, num_train_epochs=a.epochs, learning_rate=a.lr, |
| warmup_steps=a.warmup, bf16=True, eval_strategy="steps", eval_steps=a.eval_steps, |
| save_strategy="steps", save_steps=a.eval_steps, save_total_limit=a.save_total_limit, |
| load_best_model_at_end=True, metric_for_best_model="combine", greater_is_better=False, |
| logging_steps=50, gradient_checkpointing=True, |
| gradient_checkpointing_kwargs={"use_reentrant": False}, |
| dataloader_num_workers=a.num_workers, remove_unused_columns=False, report_to=[], |
| seed=a.seed, data_seed=a.seed, ignore_data_skip=True) |
|
|
| trainer = LGTrainer(model=model, args=targs, train_dataset=train_ds, eval_dataset=eval_ds, |
| data_collator=train_collator, eval_collator=eval_collator, |
| compute_metrics=metrics, preprocess_logits_for_metrics=preprocess_logits, |
| processing_class=processor) |
| print("BASELINE eval (avant fine-tune):", json.dumps(trainer.evaluate(), default=float), flush=True) |
| trainer.train() |
| trainer.save_model(os.path.join(a.out, "best")) |
| processor.save_pretrained(os.path.join(a.out, "best")) |
| print("EVAL FINALE:", json.dumps(trainer.evaluate(), default=float), flush=True) |
| print("TRAIN_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|