#!/usr/bin/env python3 """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 # w2v-bert 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) # /root/models/joint_cont_best p.add_argument("--lr", type=float, default=2e-5) # full fine-tune bas LR 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) # REUTILISE processor (vocab + FE) du checkpoint init -> tete CTC conservee 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"): # gel de la projection de features (front-end) pour stabilite/memoire 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()