waxal2026-backup / phase2_corrected /code /train_joint_aug.py
Pricile's picture
compactage apres suppression luganda
6eed659
Raw
History Blame Contribute Delete
9.34 kB
#!/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()