| |
| """LEVER C — MMS-1B adaptateurs JOINT + anti-sur-apprentissage OOD. |
| Diffs vs train_mms_adapter.py : |
| - AUGMENTATION synthetique (audiomentations) sur le TRAIN uniquement : bruit gaussien, |
| time-stretch (speed-perturb 0.9-1.1), pitch-shift, gain. => invariance locuteur/canal. |
| (100% conforme : transformations de l'audio du challenge, AUCUNE donnee externe.) |
| - SpecAugment agressif (mask_time_prob/mask_feature_prob dans la config du modele, sur GPU). |
| - eval = dev-difficile (locuteurs DISJOINTS) ; l'eval n'est PAS augmentee (collator separe). |
| - lr par defaut 1e-4 (au lieu de 1e-3) : tue le sur-apprentissage des locuteurs vus. |
| Recette adaptateurs identique : base 1B gelee, adaptateurs frais + tete CTC (~2.3M params). |
| """ |
| import argparse |
| import json |
| import os |
| import random |
| from dataclasses import dataclass |
|
|
| import jiwer |
| import numpy as np |
| import soundfile as sf |
| import torch |
| from datasets import Dataset |
| from transformers import ( |
| Trainer, |
| TrainingArguments, |
| Wav2Vec2CTCTokenizer, |
| Wav2Vec2FeatureExtractor, |
| Wav2Vec2ForCTC, |
| Wav2Vec2Processor, |
| ) |
| from transformers.trainer_pt_utils import LengthGroupedSampler |
| from audiomentations import Compose, AddGaussianNoise, TimeStretch, PitchShift, Gain |
|
|
| SR = 16000 |
| BASE = "facebook/mms-1b-all" |
| FRAMES_PER_SEC = 50 |
|
|
|
|
| class LGTrainer(Trainer): |
| """LengthGrouped sampler + collator d'eval SANS augmentation.""" |
| 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 text_key(t): |
| import unicodedata |
| t = unicodedata.normalize("NFC", t).lower() |
| return " ".join("".join(c for c in t if c.isalnum() or c.isspace()).split()) |
|
|
|
|
| def parse(): |
| p = argparse.ArgumentParser() |
| p.add_argument("--lang", default="jointC") |
| p.add_argument("--out", required=True) |
| p.add_argument("--train", nargs="+", required=True) |
| p.add_argument("--eval", required=True) |
| p.add_argument("--init", default=BASE) |
| p.add_argument("--lr", type=float, default=1e-4) |
| p.add_argument("--epochs", type=float, default=12) |
| p.add_argument("--bs", type=int, default=4) |
| p.add_argument("--grad_accum", type=int, default=16) |
| p.add_argument("--warmup", type=int, default=500) |
| p.add_argument("--eval_steps", type=int, default=300) |
| p.add_argument("--eval_subset", type=int, default=1200) |
| p.add_argument("--min_dur", type=float, default=1.5) |
| p.add_argument("--max_dur", type=float, default=30.0) |
| p.add_argument("--max_wps", type=float, default=4.0) |
| p.add_argument("--exclude_texts", default=None) |
| p.add_argument("--num_workers", type=int, default=8) |
| p.add_argument("--save_total_limit", type=int, default=4) |
| 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("--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) |
|
|
| 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)") |
| train_rows += kept |
| if a.exclude_texts: |
| banned = set() |
| for m in a.exclude_texts.split(","): |
| banned |= {text_key(r["text"]) for r in read_jsonl(m) if r["text"]} |
| before = len(train_rows) |
| train_rows = [r for r in train_rows if text_key(r["text"]) not in banned] |
| print(f"anti-fuite: {before-len(train_rows)} exclus") |
| 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[:96], 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}") |
|
|
| |
| vocab_path = os.path.join(a.out, "vocab.json") |
| chars = set() |
| for r in train_rows: |
| chars.update(r["text"]) |
| chars -= {" ", "|"} |
| vocab = {c: i for i, c in enumerate(sorted(chars))} |
| vocab["|"] = len(vocab); vocab["[UNK]"] = len(vocab); vocab["[PAD]"] = len(vocab) |
| json.dump(vocab, open(vocab_path, "w", encoding="utf-8"), ensure_ascii=False, indent=1) |
|
|
| tok = Wav2Vec2CTCTokenizer(vocab_path, unk_token="[UNK]", pad_token="[PAD]", word_delimiter_token="|") |
| fe = Wav2Vec2FeatureExtractor(feature_size=1, sampling_rate=SR, padding_value=0.0, |
| do_normalize=True, return_attention_mask=True) |
| processor = Wav2Vec2Processor(feature_extractor=fe, tokenizer=tok) |
| processor.save_pretrained(a.out) |
|
|
| 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 |
| train_rows = [r for r in train_rows if feasible(r)] |
|
|
| 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 = Wav2Vec2ForCTC.from_pretrained( |
| a.init, vocab_size=len(tok), pad_token_id=tok.pad_token_id, |
| ctc_loss_reduction="mean", ctc_zero_infinity=True, |
| attention_dropout=0.0, hidden_dropout=0.0, feat_proj_dropout=0.0, layerdrop=0.0, |
| 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, |
| ignore_mismatched_sizes=True) |
| model.init_adapter_layers() |
| model.freeze_base_model() |
| n_train = 0 |
| for name, p in model.named_parameters(): |
| if "adapter" in name.lower(): |
| p.requires_grad = True |
| if p.requires_grad: |
| n_train += p.numel() |
| print(f"params entrainables: {n_train/1e6:.2f}M (adaptateurs + lm_head) sur ~1B geles", 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=False, |
| 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) |
| 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)) |
| print("TRAIN_DONE", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|