oddadmix's picture
Upload train.py with huggingface_hub
cd20eaf verified
Raw
History Blame Contribute Delete
12.4 kB
"""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()
@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()