""" Fine-tune Whisper on Common Voice Scripted Speech 27.0 Tatar. The recipe follows the simple fine-tuning (SFT) setup of BuzzASR (arxiv:2609.09554): full fine-tuning of Whisper-large-v3, an encoder learning rate of 0.3x the decoder's, a cosine schedule, early stopping on dev, and greedy decoding with repetition penalties. Data, all drawn from the validated bucket, in one of two ways. Official split (default): - test: the official CV test split, untouched, so results compare with eval.py - dev: about 500 clips from randomly chosen light speakers, one per sentence (kept small: the test split holds most speakers, so few are left) - train: every other validated clip, sharing no speaker and no sentence with dev or test Speakers per split in CV Scripted Speech 27.0 Tatar (clips from the TSV files): split clips speakers dev 4765 9 invalidated 580 149 other 264 19 test 4983 258 train 8249 2 validated 30555 270 Test holds 258 of the 270 validated speakers, which leaves only a dozen for train and dev together. Reproduce with: dataset.groupby("split").agg( clips=("audio_path", "size"), speakers=("speaker_id", "nunique") ) Re-split (--resplit), to train on far more voices: - test: up to 40 clips each from 80 random speakers of the official test split - dev: the same from 25 other official test speakers - train: every other validated clip, sharing no speaker and no sentence with dev or test (about 165 speakers instead of 9) Test and dev come only from official test clips, which a model trained on the official split has never seen, so both kinds of model can be compared on the new test set (eval.py --test-tsv /splits/test.tsv). Labels keep their original casing and punctuation, so the model learns to write readable subtitles. Metrics are computed on normalized text, as in eval.py. Usage: python3 languages/tat/finetune.py --dry-run # show the splits, train nothing python3 languages/tat/finetune.py --resplit --dry-run python3 languages/tat/finetune.py # defaults suit an A100 80 GB python3 languages/tat/finetune.py --batch-size 2 --grad-accum 16 --optim-8bit # A100 40 GB python3 languages/tat/finetune.py --resume # continue from the last check- # point To evaluate the result on the test split: python3 languages/tat/eval.py --models languages/tat/whisper-large-v3-tt/final Requires the datacollective fork [1] with the TSV quoting fix [2], the script refuses to run if load_dataset returns fewer rows than the TSV files contain. [1] https://github.com/IlnarSelimcan/datacollective-python/ [2] https://github.com/IlnarSelimcan/datacollective-python/commit/3e9312ee2d9df9d2f4438667b66cd6b0bed0a854 """ import os import random from collections import defaultdict from collections.abc import Callable, Sequence from dataclasses import dataclass from itertools import combinations from pathlib import Path from typing import Annotated, TypedDict import librosa import numpy as np import numpy.typing as npt import pandas as pd import torch import typer from datacollective import load_dataset from transformers import ( EarlyStoppingCallback, EvalPrediction, Seq2SeqTrainer, Seq2SeqTrainingArguments, WhisperForConditionalGeneration, WhisperProcessor, WhisperTokenizer, set_seed, ) from eval import ( COMMON_VOICE_SCRIPTED_SPEECH_27_0_TATAR, SAMPLING_RATE, WHISPER_GENERATE_KWARGS, Audio, load_audio, score, ) HERE = Path(__file__).resolve().parent MODEL_ID = "openai/whisper-large-v3" OUTPUT_DIR = HERE / "whisper-large-v3-tt" LABEL_PAD = -100 """Label value ignored by the loss (PyTorch's cross-entropy default).""" SPEED_RATES = (0.9, 1.0, 1.1) """Speed perturbation factors, used only with --speed-perturb.""" # -- Types -------------------------------------------------------------------- class Example(TypedDict): """One training example: log-mel features and label token ids.""" input_features: npt.NDArray[np.float32] labels: list[int] class Batch(TypedDict): """A padded batch, as the model's forward pass expects it.""" input_features: torch.Tensor labels: torch.Tensor @dataclass(frozen=True) class Splits: """Train, dev and test clips, one row per clip.""" train: pd.DataFrame dev: pd.DataFrame test: pd.DataFrame def items(self) -> list[tuple[str, pd.DataFrame]]: """Return (name, frame) pairs in a fixed order.""" return [("train", self.train), ("dev", self.dev), ("test", self.test)] @dataclass(frozen=True) class Config: """Everything that defines one training run.""" model_id: str output_dir: Path learning_rate: float encoder_lr_ratio: float batch_size: int grad_accum: int epochs: int warmup_steps: int weight_decay: float eval_steps: int patience: int dev_clips: int max_dev_speaker_clips: int max_clips_per_speaker: int resplit: bool test_speakers: int dev_speakers: int max_held_out_speaker_clips: int speed_perturb: bool optim_8bit: bool workers: int seed: int resume: bool dry_run: bool # -- Data --------------------------------------------------------------------- def load_common_voice() -> pd.DataFrame: """ Return every CV Scripted Speech 27.0 Tatar row, after checking that none were dropped. """ if not os.getenv("MDC_API_KEY"): raise RuntimeError( "Set the MDC_API_KEY environment variable to download the dataset." ) dataset = load_dataset(COMMON_VOICE_SCRIPTED_SPEECH_27_0_TATAR) check_complete(dataset) return dataset def tsv_row_count(path: Path) -> int: """ Return the number of data rows in the TSV file at `path` (assuming it has a header line). """ with path.open(encoding="utf-8") as f: return sum(1 for _ in f) - 1 def check_complete(dataset: pd.DataFrame) -> None: """ Raise if `dataset` has fewer rows per split than its TSV file, which is what an unpatched datacollective returns. """ cv_dir = Path(dataset.audio_path.iloc[0]).parent.parent for split, loaded in dataset.split.value_counts().items(): expected = tsv_row_count(cv_dir / f"{split}.tsv") if loaded != expected: raise RuntimeError( f"load_dataset returned {loaded} '{split}' rows, but " f"{split}.tsv has {expected}. Install the patched datacollective." ) def exclude(frame: pd.DataFrame, held_out: pd.DataFrame) -> pd.DataFrame: """ Return the rows of `frame` sharing no speaker and no sentence with `held_out`. """ shared = frame.speaker_id.isin( held_out.speaker_id ) | frame.sentence_id.isin(held_out.sentence_id) return frame[~shared] def pick_dev_speakers( pool: pd.DataFrame, target_clips: int, max_clips: int, seed: int ) -> set[str]: """ Return random "light" speakers from `pool` until together they have `target_clips` clips. Contributions are very uneven: a couple of speakers recorded thousands of clips each, most others a few dozen. Only speakers with at most `max_clips` clips qualify ("light"), because every dev speaker's clips are removed from train to keep the two splits speaker-disjoint. A heavy speaker in dev would take thousands of clips out of an already small training set; a light one costs only a few dozen. """ counts = pool.speaker_id.value_counts() light = counts[counts <= max_clips].sample(frac=1, random_state=seed) clips_before = light.cumsum().shift(fill_value=0) return set(light[clips_before < target_clips].index) def cap_per_speaker(frame: pd.DataFrame, cap: int, seed: int) -> pd.DataFrame: """Return at most `cap` random clips per speaker from `frame`.""" shuffled = frame.sample(frac=1, random_state=seed) return shuffled.groupby("speaker_id").head(cap) def check_disjoint(splits: Splits) -> None: """Raise if any two splits share a speaker or a sentence.""" for (name_a, a), (name_b, b) in combinations(splits.items(), 2): for column in ("speaker_id", "sentence_id"): shared = set(a[column]) & set(b[column]) if shared: raise ValueError( f"{name_a} and {name_b} share {len(shared)} values of {column}" ) def make_splits( dataset: pd.DataFrame, dev_clips: int, max_dev_speaker_clips: int, max_clips_per_speaker: int, seed: int, ) -> Splits: """ Return train/dev/test splits: the official test split, a small dev set of light speakers, and every other validated clip as train. `max_clips_per_speaker` of 0 means no cap. """ test = dataset[dataset.split == "test"] pool = exclude(dataset[dataset.split == "validated"], test) speakers = pick_dev_speakers(pool, dev_clips, max_dev_speaker_clips, seed) if not speakers: raise ValueError( f"No speaker outside test has at most {max_dev_speaker_clips} clips, " "so dev would be empty; raise --max-dev-speaker-clips." ) dev = pool[pool.speaker_id.isin(speakers)].drop_duplicates("sentence_id") train = exclude(pool, dev) if max_clips_per_speaker: train = cap_per_speaker(train, max_clips_per_speaker, seed) splits = Splits( train=train.sample(frac=1, random_state=seed).reset_index(drop=True), dev=dev.reset_index(drop=True), test=test.reset_index(drop=True), ) check_disjoint(splits) return splits def sample_speakers( frame: pd.DataFrame, n: int, cap: int, seed: int ) -> pd.DataFrame: """ Return up to `cap` clips from each of `n` random speakers in `frame`, one clip per sentence. """ speakers = frame.speaker_id.drop_duplicates() chosen = speakers.sample(n=min(n, len(speakers)), random_state=seed) clips = frame[frame.speaker_id.isin(chosen)].drop_duplicates("sentence_id") return cap_per_speaker(clips, cap, seed) def make_resplits( dataset: pd.DataFrame, test_speakers: int, dev_speakers: int, max_held_out_speaker_clips: int, max_clips_per_speaker: int, seed: int, ) -> Splits: """ Return train/dev/test splits drawn afresh: test and dev from speakers of the official test split, train from every other validated clip. `max_clips_per_speaker` of 0 means no cap on train. """ official_test = dataset[dataset.split == "test"] test = sample_speakers( official_test, test_speakers, max_held_out_speaker_clips, seed ) dev = sample_speakers( exclude(official_test, test), dev_speakers, max_held_out_speaker_clips, seed, ) train = exclude(exclude(dataset[dataset.split == "validated"], test), dev) if max_clips_per_speaker: train = cap_per_speaker(train, max_clips_per_speaker, seed) splits = Splits( train=train.sample(frac=1, random_state=seed).reset_index(drop=True), dev=dev.reset_index(drop=True), test=test.reset_index(drop=True), ) check_disjoint(splits) return splits def build_splits(dataset: pd.DataFrame, cfg: Config) -> Splits: """Return the splits `cfg` asks for: official (default) or re-split.""" if cfg.resplit: return make_resplits( dataset, cfg.test_speakers, cfg.dev_speakers, cfg.max_held_out_speaker_clips, cfg.max_clips_per_speaker, cfg.seed, ) return make_splits( dataset, cfg.dev_clips, cfg.max_dev_speaker_clips, cfg.max_clips_per_speaker, cfg.seed, ) def describe(splits: Splits) -> pd.DataFrame: """Return clip, speaker and sentence counts per split.""" return pd.DataFrame.from_dict( { name: { "clips": len(frame), "speakers": frame.speaker_id.nunique(), "sentences": frame.sentence_id.nunique(), "top_speaker_share": round( frame.speaker_id.value_counts(normalize=True).iloc[0], 3 ), } for name, frame in splits.items() }, orient="index", ) def save_splits(splits: Splits, directory: Path) -> None: """ Write each split's clips to `directory`/.tsv, for reproducibility. """ directory.mkdir(parents=True, exist_ok=True) columns = ["audio_path", "transcription", "speaker_id", "sentence_id"] for name, frame in splits.items(): frame[columns].to_csv(directory / f"{name}.tsv", sep="\t", index=False) # -- Features ----------------------------------------------------------------- def change_speed(audio: Audio, rate: float) -> Audio: """ Return `audio` played `rate` times faster, pitch included, as in Kaldi-style speed perturbation. """ if rate == 1.0: return audio return librosa.resample( audio, orig_sr=round(SAMPLING_RATE * rate), target_sr=SAMPLING_RATE ) class WhisperDataset: """Map-style dataset that decodes audio and tokenizes text on access.""" def __init__( self, frame: pd.DataFrame, processor: WhisperProcessor, speed_rates: Sequence[float] = (), ) -> None: self.paths = frame.audio_path.tolist() self.texts = frame.transcription.tolist() self.processor = processor self.speed_rates = speed_rates def __len__(self) -> int: return len(self.paths) def __getitem__(self, i: int) -> Example: audio = load_audio(self.paths[i]) if self.speed_rates: # `random` is reseeded per DataLoader worker, unlike numpy's global # RNG audio = change_speed(audio, random.choice(self.speed_rates)) features = self.processor.feature_extractor( audio, sampling_rate=SAMPLING_RATE ).input_features[0] labels = self.processor.tokenizer(self.texts[i]).input_ids return {"input_features": features, "labels": labels} @dataclass(frozen=True) class Collator: """ Pads examples into a batch; padded label positions are ignored by the loss. """ processor: WhisperProcessor decoder_start_token_id: int def __call__(self, examples: list[Example]) -> Batch: features = self.processor.feature_extractor.pad( [{"input_features": e["input_features"]} for e in examples], return_tensors="pt", ) labels = self.processor.tokenizer.pad( [{"input_ids": e["labels"]} for e in examples], return_tensors="pt" ) ids = labels["input_ids"].masked_fill( labels.attention_mask.ne(1), LABEL_PAD ) # the model prepends the start token itself during training if (ids[:, 0] == self.decoder_start_token_id).all(): ids = ids[:, 1:] return {"input_features": features["input_features"], "labels": ids} # -- Model -------------------------------------------------------------------- def load_model( model_id: str, ) -> tuple[WhisperForConditionalGeneration, WhisperProcessor]: """Return the model and processor, set up to transcribe Tatar.""" processor = WhisperProcessor.from_pretrained( model_id, language="tatar", task="transcribe" ) model = WhisperForConditionalGeneration.from_pretrained( model_id, dtype=torch.float32, # transformers 5 defaults to float16 ) model.config.use_cache = False # incompatible with gradient checkpointing configure_generation(model) return model, processor def configure_generation(model: WhisperForConditionalGeneration) -> None: """ Decode Tatar with the same greedy settings as eval.py, both for dev evaluation and in the saved model. """ for key, value in WHISPER_GENERATE_KWARGS.items(): setattr(model.generation_config, key, value) model.generation_config.forced_decoder_ids = None model.config.forced_decoder_ids = None def parameter_groups( model: torch.nn.Module, lr: float, encoder_lr_ratio: float, weight_decay: float, ) -> list[dict[str, object]]: """ Return AdamW parameter groups: the encoder gets `encoder_lr_ratio` times the decoder's learning rate; biases and norms get no weight decay. """ grouped: dict[tuple[bool, bool], list[torch.nn.Parameter]] = defaultdict( list ) for name, param in model.named_parameters(): if param.requires_grad: is_encoder = name.startswith("model.encoder.") decays = param.ndim >= 2 grouped[(is_encoder, decays)].append(param) return [ { "params": params, "lr": lr * encoder_lr_ratio if is_encoder else lr, "weight_decay": weight_decay if decays else 0.0, } for (is_encoder, decays), params in grouped.items() ] def build_optimizer( model: torch.nn.Module, cfg: Config ) -> torch.optim.Optimizer: """Return AdamW, or its 8-bit version, over the model's parameter groups.""" groups = parameter_groups( model, cfg.learning_rate, cfg.encoder_lr_ratio, cfg.weight_decay ) if cfg.optim_8bit: import bitsandbytes as bnb return bnb.optim.AdamW8bit(groups) return torch.optim.AdamW(groups) # -- Training ------------------------------------------------------------------ def make_compute_metrics( tokenizer: WhisperTokenizer, ) -> Callable[[EvalPrediction], dict[str, float]]: """Return a function scoring generated predictions with eval.score.""" def decode(ids: npt.NDArray[np.int64]) -> list[str]: ids = np.where(ids == LABEL_PAD, tokenizer.pad_token_id, ids) return tokenizer.batch_decode(ids, skip_special_tokens=True) def compute_metrics(pred: EvalPrediction) -> dict[str, float]: texts = pd.DataFrame( { "reference": decode(pred.label_ids), "hypothesis": decode(pred.predictions), } ) return score(texts) return compute_metrics def use_bf16() -> bool: """ Return whether the GPU supports bfloat16, which is safer than fp16 for training. """ return torch.cuda.is_available() and torch.cuda.is_bf16_supported() def training_args(cfg: Config) -> Seq2SeqTrainingArguments: """Return the Trainer settings for `cfg`.""" return Seq2SeqTrainingArguments( output_dir=str(cfg.output_dir), per_device_train_batch_size=cfg.batch_size, per_device_eval_batch_size=cfg.batch_size * 4, gradient_accumulation_steps=cfg.grad_accum, learning_rate=cfg.learning_rate, # scheduler scales each group's own rate lr_scheduler_type="cosine", warmup_steps=cfg.warmup_steps, num_train_epochs=cfg.epochs, bf16=use_bf16(), gradient_checkpointing=True, gradient_checkpointing_kwargs={"use_reentrant": False}, eval_strategy="steps", eval_steps=cfg.eval_steps, save_strategy="steps", save_steps=cfg.eval_steps, save_total_limit=2, load_best_model_at_end=True, metric_for_best_model="cer", greater_is_better=False, predict_with_generate=True, generation_max_length=225, logging_steps=25, report_to=["tensorboard"], dataloader_num_workers=cfg.workers, remove_unused_columns=False, label_names=["labels"], seed=cfg.seed, ) def has_checkpoint(directory: Path) -> bool: """Return whether `directory` contains a Trainer checkpoint.""" return directory.is_dir() and any(directory.glob("checkpoint-*")) def build_trainer( cfg: Config, splits: Splits ) -> tuple[Seq2SeqTrainer, WhisperProcessor]: """Return a Trainer for `cfg` on `splits`, and the processor it uses.""" model, processor = load_model(cfg.model_id) speed_rates = SPEED_RATES if cfg.speed_perturb else () trainer = Seq2SeqTrainer( model=model, args=training_args(cfg), train_dataset=WhisperDataset(splits.train, processor, speed_rates), eval_dataset=WhisperDataset(splits.dev, processor), data_collator=Collator(processor, model.config.decoder_start_token_id), compute_metrics=make_compute_metrics(processor.tokenizer), processing_class=processor, optimizers=(build_optimizer(model, cfg), None), callbacks=[ EarlyStoppingCallback(early_stopping_patience=cfg.patience) ], ) return trainer, processor def run(cfg: Config) -> Path | None: """Train according to `cfg`; return the final model's directory, or None on a dry run.""" set_seed(cfg.seed) splits = build_splits(load_common_voice(), cfg) print(describe(splits).to_string()) if cfg.dry_run: return None save_splits(splits, cfg.output_dir / "splits") trainer, processor = build_trainer(cfg, splits) trainer.train( resume_from_checkpoint=cfg.resume and has_checkpoint(cfg.output_dir) ) final = cfg.output_dir / "final" trainer.save_model( str(final) ) # the best checkpoint, thanks to load_best_model_at_end processor.save_pretrained(str(final)) print(f"Best model saved to {final}") return final # -- Main --------------------------------------------------------------------- def main( model_id: Annotated[ str, typer.Option(help="Model to fine-tune.") ] = MODEL_ID, output_dir: Annotated[ Path, typer.Option(help="Checkpoints, splits and final model.") ] = OUTPUT_DIR, learning_rate: Annotated[ float, typer.Option( help="Decoder learning rate (BuzzASR: 3.24e-6, 6.48e-6, 2e-5)." ), ] = 6.48e-6, encoder_lr_ratio: Annotated[ float, typer.Option( help="Encoder learning rate as a fraction of the decoder's." ), ] = 0.3, batch_size: Annotated[int, typer.Option(help="Clips per GPU step.")] = 4, grad_accum: Annotated[ int, typer.Option( help="Steps per optimizer update; effective batch = batch size x this." ), ] = 8, epochs: Annotated[int, typer.Option(help="Maximum epochs.")] = 6, warmup_steps: Annotated[ int, typer.Option(help="Learning-rate warmup steps.") ] = 150, weight_decay: Annotated[ float, typer.Option(help="AdamW weight decay.") ] = 4e-5, eval_steps: Annotated[ int, typer.Option(help="Evaluate and save every N updates.") ] = 100, patience: Annotated[ int, typer.Option( help="Stop after N evaluations without dev CER improvement." ), ] = 3, dev_clips: Annotated[ int, typer.Option(help="Approximate dev set size, in clips.") ] = 500, max_dev_speaker_clips: Annotated[ int, typer.Option( help="Only speakers with at most this many clips go to dev." ), ] = 200, max_clips_per_speaker: Annotated[ int, typer.Option(help="Cap on train clips per speaker (0 = no cap).") ] = 0, resplit: Annotated[ bool, typer.Option(help="Draw new test/dev/train splits (see docstring)."), ] = False, test_speakers: Annotated[ int, typer.Option(help="With --resplit: speakers in the test set.") ] = 80, dev_speakers: Annotated[ int, typer.Option(help="With --resplit: speakers in the dev set.") ] = 25, max_held_out_speaker_clips: Annotated[ int, typer.Option( help="With --resplit: clips per test/dev speaker, at most." ), ] = 40, speed_perturb: Annotated[ bool, typer.Option(help="Randomly change speed by 0.9x/1.1x in training."), ] = False, optim_8bit: Annotated[ bool, typer.Option( help="Use 8-bit AdamW (bitsandbytes) to save GPU memory." ), ] = False, workers: Annotated[ int, typer.Option(help="DataLoader worker processes.") ] = 4, seed: Annotated[ int, typer.Option(help="Random seed for splits and training.") ] = 42, resume: Annotated[ bool, typer.Option(help="Resume from the last checkpoint.") ] = False, dry_run: Annotated[ bool, typer.Option(help="Only build and describe the splits.") ] = False, ) -> Path | None: """Fine-tune Whisper on Common Voice 27.0 Tatar.""" return run( Config(**locals()) ) # the parameters are exactly Config's fields if __name__ == "__main__": typer.run(main)