"""Full fine-tune of Whisper-small on the oddadmix multi-dialect Arabic set. * Base: openai/whisper-small (244M), full fine-tune (no LoRA). * Data: oddadmix/dialectal-arabic-lahgtna-v2-smaller-augmented (already 16 kHz mono; has an `augmentation` column). * Targets: cleaned with normalize.clean_text (tashkil + tags stripped). * Features: log-mel computed ON THE FLY in the collator, so we never materialize ~38 GB of cached features to disk. * Precision: bf16 on a 32 GB GPU. Auth: the dataset is private. Export a token first: export HF_TOKEN=hf_xxx # account with access to oddadmix/... Run: uv run accelerate launch train.py # or: uv run python train.py """ from __future__ import annotations import argparse import json from dataclasses import dataclass from datetime import datetime, timezone from pathlib import Path from typing import Any import torch from datasets import load_dataset import jiwer from transformers import ( WhisperProcessor, WhisperForConditionalGeneration, Seq2SeqTrainer, Seq2SeqTrainingArguments, ) from normalize import clean_text MAX_AUDIO_SECONDS = 30.0 # Whisper encoder hard limit MIN_AUDIO_SECONDS = 0.5 MAX_LABEL_TOKENS = 448 # Whisper decoder max target length def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser() p.add_argument("--base_model", default="openai/whisper-small") p.add_argument("--dataset", default="oddadmix/dialectal-arabic-lahgtna-v2-smaller-augmented") p.add_argument("--language", default="ar") p.add_argument("--run_name", default=None, help="Name for this run. Outputs go to runs//. " "Defaults to the base model name.") p.add_argument("--output_dir", default=None, help="Override output dir (default: runs/).") p.add_argument("--notes", default="", help="Free-text note recorded in the run summary / README.") p.add_argument("--per_device_train_batch_size", type=int, default=16) p.add_argument("--per_device_eval_batch_size", type=int, default=8) p.add_argument("--gradient_accumulation_steps", type=int, default=2) p.add_argument("--learning_rate", type=float, default=1e-5) p.add_argument("--warmup_steps", type=int, default=500) p.add_argument("--max_steps", type=int, default=6000) p.add_argument("--eval_steps", type=int, default=500) p.add_argument("--save_steps", type=int, default=500) p.add_argument("--num_workers", type=int, default=8) p.add_argument("--normalize_letters", action="store_true", help="Also fold أإآ->ا, ى->ي, ة->ه (off by default).") p.add_argument("--eval_only_frac", type=float, default=1.0, help="Use a fraction of the test split for periodic eval (speed).") p.add_argument("--resume_from_checkpoint", default=None, help="Path to a checkpoint dir to resume training from.") return p.parse_args() @dataclass class DataCollator: """Extract log-mel features from raw audio and tokenize cleaned labels.""" processor: WhisperProcessor normalize_letters: bool decoder_start_token_id: int def __call__(self, batch: list[dict[str, Any]]) -> dict[str, torch.Tensor]: fe = self.processor.feature_extractor tok = self.processor.tokenizer arrays = [ex["audio"]["array"] for ex in batch] feats = fe(arrays, sampling_rate=16000, return_tensors="pt") out = {"input_features": feats.input_features} texts = [clean_text(ex["text"], self.normalize_letters) for ex in batch] # tok(texts) prepends the Whisper prefix (<|sot|><|ar|><|transcribe|> # <|notimestamps|>) and appends <|eot|>. The Trainer re-prepends the # decoder-start token, so we strip the leading <|sot|> below. # truncation is a safety net; over-length rows are filtered out in main() label_ids = tok(texts, max_length=MAX_LABEL_TOKENS, truncation=True).input_ids labels = tok.pad({"input_ids": label_ids}, return_tensors="pt") # mask padding so it's ignored by the loss mask = labels.attention_mask.ne(1) labels_ids = labels.input_ids.masked_fill(mask, -100) # The tokenizer prepends <|startoftranscript|> (== decoder_start_token_id). # The Trainer re-prepends it when building decoder_input_ids, so strip it # here to avoid a doubled start token. NOTE: Whisper's bos_token_id is # <|endoftext|> (50257), NOT the sot token, so we must compare against # decoder_start_token_id explicitly. if (labels_ids[:, 0] == self.decoder_start_token_id).all().cpu().item(): labels_ids = labels_ids[:, 1:] out["labels"] = labels_ids return out def build_metrics(processor, normalize_letters): tok = processor.tokenizer def compute_metrics(pred): pred_ids = pred.predictions label_ids = pred.label_ids label_ids[label_ids == -100] = tok.pad_token_id pred_str = tok.batch_decode(pred_ids, skip_special_tokens=True) ref_str = tok.batch_decode(label_ids, skip_special_tokens=True) preds = [clean_text(p, normalize_letters) for p in pred_str] refs = [clean_text(r, normalize_letters) for r in ref_str] # jiwer needs non-empty references; drop any degenerate pairs pairs = [(p, r) for p, r in zip(preds, refs) if r.strip()] if not pairs: return {"wer": 1.0, "cer": 1.0} preds, refs = map(list, zip(*pairs)) return { "wer": jiwer.wer(refs, preds), "cer": jiwer.cer(refs, preds), } return compute_metrics def keep_row(text: str, duration: float, normalize_letters: bool) -> bool: if duration is None or not (MIN_AUDIO_SECONDS <= duration <= MAX_AUDIO_SECONDS): return False return bool(clean_text(text, normalize_letters)) def main() -> None: args = parse_args() run_name = args.run_name or Path(args.base_model).name output_dir = args.output_dir or f"runs/{run_name}" Path(output_dir).mkdir(parents=True, exist_ok=True) print(f"run '{run_name}' -> {output_dir}") processor = WhisperProcessor.from_pretrained( args.base_model, language=args.language, task="transcribe" ) ds = load_dataset(args.dataset) # Filter on text+duration only -> no audio decode during filtering. nl = args.normalize_letters ds = ds.filter( lambda text, duration: keep_row(text, duration, nl), input_columns=["text", "duration"], num_proc=args.num_workers, ) # Drop rows whose (cleaned) transcript exceeds Whisper's 448-token target # limit. These are corrupt/mislabeled clips (e.g. a huge text blob on a # short clip) and would crash training. Tokenizing text is cheap; keep it # single-process to avoid fast-tokenizer fork deadlocks. tok = processor.tokenizer before = {split: len(ds[split]) for split in ds} ds = ds.filter( lambda text: len(tok(clean_text(text, nl)).input_ids) <= MAX_LABEL_TOKENS, input_columns=["text"], num_proc=1, ) after = {split: len(ds[split]) for split in ds} print(f"kept {after} (dropped over-length: " f"{ {s: before[s] - after[s] for s in before} })") eval_ds = ds["test"] if args.eval_only_frac < 1.0: n = max(1, int(len(eval_ds) * args.eval_only_frac)) eval_ds = eval_ds.select(range(n)) # Force fp32 load: some checkpoints (e.g. large-v3-turbo) ship in fp16, which # crashes generate() at eval (no autocast -> fp16 weights vs fp32 features). # bf16-mixed training autocasts from fp32 weights, same as small/medium. model = WhisperForConditionalGeneration.from_pretrained( args.base_model, torch_dtype=torch.float32) # Standard Whisper fine-tuning: let the model learn the language, don't # force decoder ids or suppress tokens during training. model.config.forced_decoder_ids = None model.config.suppress_tokens = [] model.generation_config.language = args.language model.generation_config.task = "transcribe" model.generation_config.forced_decoder_ids = None model.config.use_cache = False # required with gradient checkpointing collator = DataCollator( processor=processor, normalize_letters=nl, decoder_start_token_id=model.config.decoder_start_token_id, ) training_args = Seq2SeqTrainingArguments( output_dir=output_dir, run_name=run_name, per_device_train_batch_size=args.per_device_train_batch_size, per_device_eval_batch_size=args.per_device_eval_batch_size, gradient_accumulation_steps=args.gradient_accumulation_steps, learning_rate=args.learning_rate, warmup_steps=args.warmup_steps, max_steps=args.max_steps, gradient_checkpointing=True, bf16=True, fp16=False, eval_strategy="steps", eval_steps=args.eval_steps, save_steps=args.save_steps, logging_steps=25, report_to=["tensorboard"], predict_with_generate=True, generation_max_length=225, save_total_limit=3, load_best_model_at_end=True, metric_for_best_model="wer", greater_is_better=False, dataloader_num_workers=args.num_workers, remove_unused_columns=False, # collator needs the raw 'audio' column label_names=["labels"], ) trainer = Seq2SeqTrainer( model=model, args=training_args, train_dataset=ds["train"], eval_dataset=eval_ds, data_collator=collator, compute_metrics=build_metrics(processor, nl), processing_class=processor, ) trainer.train(resume_from_checkpoint=args.resume_from_checkpoint) trainer.save_model(output_dir) processor.save_pretrained(output_dir) # --- versioned run summary ------------------------------------------- evals = [h for h in trainer.state.log_history if "eval_wer" in h] best = min(evals, key=lambda h: h["eval_wer"]) if evals else {} train_logs = [h for h in trainer.state.log_history if "train_runtime" in h] summary = { "run_name": run_name, "base_model": args.base_model, "dataset": args.dataset, "language": args.language, "output_dir": output_dir, "notes": args.notes, "hyperparams": { "learning_rate": args.learning_rate, "warmup_steps": args.warmup_steps, "max_steps": args.max_steps, "per_device_train_batch_size": args.per_device_train_batch_size, "gradient_accumulation_steps": args.gradient_accumulation_steps, "effective_batch_size": args.per_device_train_batch_size * args.gradient_accumulation_steps, "normalize_letters": args.normalize_letters, }, "train_examples": len(ds["train"]), "eval_examples": len(eval_ds), "best_wer": round(best.get("eval_wer", float("nan")), 4), "best_cer": round(best.get("eval_cer", float("nan")), 4), "best_step": best.get("step"), "best_epoch": round(best.get("epoch", 0), 2), "best_checkpoint": trainer.state.best_model_checkpoint, "train_runtime_sec": round(train_logs[-1]["train_runtime"]) if train_logs else None, "finished_at": datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M UTC"), "eval_history": [ {"step": h["step"], "wer": round(h["eval_wer"], 4), "cer": round(h["eval_cer"], 4)} for h in evals ], } summary_path = Path(output_dir) / "summary.json" summary_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False)) print(f"Done. Model saved to {output_dir}") print(f"Best WER {summary['best_wer']} / CER {summary['best_cer']} " f"@ step {summary['best_step']}. Summary -> {summary_path}") # refresh the README run log (best-effort) try: import log_runs log_runs.update_readme() print("README run log updated.") except Exception as e: # never fail training over docs print(f"(README auto-update skipped: {e})") if __name__ == "__main__": main()