| |
| """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 |
|
|
|
|
| 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) |
| 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_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 |
|
|
| |
| 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) |
| |
| 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 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() |
|
|