Automatic Speech Recognition
Transformers
Safetensors
Arabic
whisper
arabic
dialectal-arabic
asr
Eval Results (legacy)
Instructions to use oddadmix/whisper-medium-arabic-dialectal with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use oddadmix/whisper-medium-arabic-dialectal with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("automatic-speech-recognition", model="oddadmix/whisper-medium-arabic-dialectal")# Load model directly from transformers import AutoProcessor, AutoModelForSpeechSeq2Seq processor = AutoProcessor.from_pretrained("oddadmix/whisper-medium-arabic-dialectal") model = AutoModelForSpeechSeq2Seq.from_pretrained("oddadmix/whisper-medium-arabic-dialectal", device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """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/<run_name>/. " | |
| "Defaults to the base model name.") | |
| p.add_argument("--output_dir", default=None, | |
| help="Override output dir (default: runs/<run_name>).") | |
| 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() | |
| 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() | |