""" Fine-tune Whisper-tiny on minds14 English (ASR). Deterministic split: first N examples for train, rest for test. """ # ========================= # Config (edit here) # ========================= MODEL_NAME = "openai/whisper-tiny" SLICE_TRAIN_N = 450 # first N examples for training, rest for test MAX_INPUT_LENGTH = 30.0 # max audio length (seconds) OUTPUT_DIR = "./whisper-tiny-en-US" LANGUAGE = "english" # Training params LEARNING_RATE = 2e-5 LR_SCHEDULER_TYPE = "cosine" WARMUP_STEPS = 50 EPOCHS = 15 PER_DEVICE_TRAIN_BATCH = 32 PER_DEVICE_EVAL_BATCH = 32 GRAD_ACCUMULATION = 1 GRAD_CHECKPOINT = True GRAD_CLIP_NORM = 1.0 WEIGHT_DECAY = 0.01 # Precision (keep False on MPS/CPU) DEVICE = "mps" # "auto" | "cpu" | "cuda" | "mps" FP16 = DEVICE == "cuda" # only CUDA supports FP16 FP16_FULL_EVAL = DEVICE == "cuda" # only CUDA supports FP16 BF16 = DEVICE == "cuda" # only CUDA supports BF16 # Eval / logging EVAL_STRATEGY= "epoch" # "steps" or "epoch" LOGGING_STEPS= 5 GEN_MAX_LENGTH = 225 PREDICT_WITH_GEN = True LOAD_BEST_MODEL = True METRIC_BEST = "wer" GREATER_IS_BETTER = False PUSH_TO_HUB = True # ========================= # Imports # ========================= import torch import numpy as np import evaluate from typing import Any, Dict, List, Union from datasets import load_dataset, DatasetDict, Audio from dataclasses import dataclass from transformers import ( WhisperProcessor, WhisperForConditionalGeneration, Seq2SeqTrainingArguments, Seq2SeqTrainer, ) from transformers.models.whisper.english_normalizer import BasicTextNormalizer # Device resolution def resolve_device(device_pref: str): if device_pref == "auto": if torch.cuda.is_available(): return torch.device("cuda") elif torch.backends.mps.is_available(): return torch.device("mps") else: return torch.device("cpu") else: return torch.device(device_pref) device = resolve_device(DEVICE) print(f"Using device: {device}") # ========================= # Load dataset & split # ========================= minds = load_dataset("PolyAI/minds14", "en-US") train_dataset = minds["train"].select(range(SLICE_TRAIN_N)) test_dataset = minds["train"].select(range(SLICE_TRAIN_N, len(minds["train"]))) minds = DatasetDict({"train": train_dataset, "test": test_dataset}) minds = minds.select_columns(["audio", "transcription"]) # ========================= # Processor & audio casting # ========================= processor = WhisperProcessor.from_pretrained(MODEL_NAME, language=LANGUAGE, task="transcribe") sampling_rate = processor.feature_extractor.sampling_rate minds = minds.cast_column("audio", Audio(sampling_rate=sampling_rate)) # ========================= # Preprocessing # ========================= def prepare_dataset(example): audio = example["audio"] out = processor( audio=audio["array"], sampling_rate=audio["sampling_rate"], text=example["transcription"], ) out["input_length"] = len(audio["array"]) / audio["sampling_rate"] return out remove_cols = minds["train"].column_names minds = minds.map(prepare_dataset, remove_columns=remove_cols, num_proc=1) # ========================= # Filtering # ========================= def is_audio_in_length_range(length): return length < MAX_INPUT_LENGTH for split in minds.keys(): minds[split] = minds[split].filter(is_audio_in_length_range, input_columns=["input_length"]) # ========================= # Data collator # ========================= @dataclass class DataCollatorSpeechSeq2SeqWithPadding: processor: Any def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]: input_features = [{"input_features": f["input_features"][0]} for f in features] batch = self.processor.feature_extractor.pad(input_features, return_tensors="pt") label_features = [{"input_ids": f["labels"]} for f in features] labels_batch = self.processor.tokenizer.pad(label_features, return_tensors="pt") labels = labels_batch["input_ids"].masked_fill(labels_batch.attention_mask.ne(1), -100) if (labels[:, 0] == self.processor.tokenizer.bos_token_id).all().cpu().item(): labels = labels[:, 1:] batch["labels"] = labels return batch data_collator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor) # ========================= # Metrics (WER) # ========================= metric = evaluate.load("wer") normalizer = BasicTextNormalizer() def compute_metrics(pred): pred_ids = pred.predictions[0] if isinstance(pred.predictions, (tuple, list)) else pred.predictions label_ids = np.array(pred.label_ids, copy=True) pad_id = processor.tokenizer.pad_token_id label_ids = np.where(label_ids != -100, label_ids, pad_id) pred_str = processor.batch_decode(pred_ids, skip_special_tokens=True) label_str = processor.batch_decode(label_ids, skip_special_tokens=True) wer_ortho = metric.compute(predictions=pred_str, references=label_str) pred_str_norm = [normalizer(s) for s in pred_str] label_str_norm = [normalizer(s) for s in label_str] keep = [i for i, s in enumerate(label_str_norm) if len(s) > 0] pred_str_norm = [pred_str_norm[i] for i in keep] label_str_norm = [label_str_norm[i] for i in keep] wer = metric.compute(predictions=pred_str_norm, references=label_str_norm) return {"wer_ortho": wer_ortho, "wer": wer} # ========================= # Model # ========================= model = WhisperForConditionalGeneration.from_pretrained(MODEL_NAME) model.to(device) model.config.use_cache = False forced_ids = processor.get_decoder_prompt_ids(language=LANGUAGE, task="transcribe") model.generation_config.forced_decoder_ids = forced_ids model.generation_config.use_cache = True # ========================= # Training args # ========================= training_args = Seq2SeqTrainingArguments( # Logging output_dir=OUTPUT_DIR, logging_steps=LOGGING_STEPS, report_to=["tensorboard"], # Optimization per_device_train_batch_size=PER_DEVICE_TRAIN_BATCH, gradient_accumulation_steps=GRAD_ACCUMULATION, max_grad_norm=GRAD_CLIP_NORM, optim="adamw_torch", adam_beta1=0.9, adam_beta2=0.999, adam_epsilon=1e-08, # Regulation learning_rate=LEARNING_RATE, lr_scheduler_type=LR_SCHEDULER_TYPE, warmup_steps=WARMUP_STEPS, weight_decay=WEIGHT_DECAY, # Training length num_train_epochs=EPOCHS, # Checkpointing / Best model save_strategy=EVAL_STRATEGY, load_best_model_at_end=LOAD_BEST_MODEL, metric_for_best_model=METRIC_BEST, # e.g. "wer" greater_is_better=GREATER_IS_BETTER, save_total_limit=3, # Precision / Memory gradient_checkpointing=GRAD_CHECKPOINT, fp16=FP16, fp16_full_eval=FP16_FULL_EVAL, bf16=BF16, # Evaluation eval_strategy=EVAL_STRATEGY, per_device_eval_batch_size=PER_DEVICE_EVAL_BATCH, predict_with_generate=PREDICT_WITH_GEN, # Generation for eval generation_max_length=GEN_MAX_LENGTH, ) # ========================= # Trainer # ========================= trainer = Seq2SeqTrainer( args=training_args, model=model, train_dataset=minds["train"], eval_dataset=minds["test"], data_collator=data_collator, compute_metrics=compute_metrics, tokenizer=processor, ) # ========================= # Train # ========================= trainer.train() # ========================= # Final evaluation # ========================= print("\n=== Final evaluation (TRAIN) ===") print(trainer.evaluate(eval_dataset=minds["train"])) print("\n=== Final evaluation (TEST) ===") print(trainer.evaluate(eval_dataset=minds["test"])) # ========================= # Push to hub # ========================= if PUSH_TO_HUB: kwargs = { "dataset_tags": "PolyAI/minds14", "finetuned_from": MODEL_NAME, "tasks": "automatic-speech-recognition", } trainer.push_to_hub(**kwargs)