File size: 11,004 Bytes
6eed659
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
#!/usr/bin/env python3
"""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):
        # bascule temporairement sur le collator sans aug pour construire le dataloader d'eval
        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)      # anti-surapprentissage
    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 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)

    # ---------- AUGMENTATION synthetique (train only) ----------
    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  # jamais ignorer un vrai echec, mais une aug ratee ne casse pas le batch
                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)

    # ---------- MODELE MMS a ADAPTATEURS + SpecAugment agressif ----------
    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()