#!/usr/bin/env python3 """MMS-1B-all fine-tuning à ADAPTATEURS (recette officielle MMS, benchmark arXiv:2512.10968). On gèle le modèle de base (1B params) et on n'entraîne que les adaptateurs de langue (~2,5M params) + la tête CTC. C'est ce qui rend MMS stable — le CTC brut sur mms-300m collapse (tout-blank), alors que l'adaptateur converge proprement. Diffs vs train_xlsr.py : init=facebook/mms-1b-all ; init_adapter_layers()+freeze_base_model()+ dégel des adaptateurs ; lr plus élevé (1e-3, seuls les adaptateurs bougent) ; dropouts à 0. """ 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 SR = 16000 BASE = "facebook/mms-1b-all" FRAMES_PER_SEC = 50 # wav2vec2-large, downsample 320x ; les adaptateurs ne changent pas le frame-rate class LGTrainer(Trainer): 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 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="lin") 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-3) # adaptateurs -> lr élevé OK p.add_argument("--epochs", type=float, default=8) 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=400) p.add_argument("--eval_subset", type=int, default=800) 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=7) p.add_argument("--save_total_limit", type=int, default=2) 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)}") # vocab char construit sur le train 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) @dataclass class Collator: def __call__(self, feats): audio = [sf.read(f["audio"], dtype="float32")[0] for f in feats] 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 # ---------- MODELE MMS à ADAPTATEURS ---------- 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, ignore_mismatched_sizes=True) # adaptateurs frais pour notre vocab + gel du modele de base model.init_adapter_layers() model.freeze_base_model() # degeler UNIQUEMENT les adaptateurs (freeze_base_model garde lm_head entrainable) 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 gelés", 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=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()