Hello! Welcome back.
In our last lesson, we delved into the architecture and training strategy of OpenAI's Whisper. We saw that its remarkable zero-shot performance stems from being pre-trained on a massive and diverse dataset of 680,000 hours of weakly supervised audio. This makes Whisper a powerful generalist for speech recognition.
However, what if you need a specialist? What if your application involves specific technical jargon, unique accents, or a low-resource language that Whisper struggles with? This is where fine-tuning comes in.
This lesson directly addresses the learning outcome: Fine-tune a pretrained Whisper model on a custom speech dataset using the Hugging Face ecosystem. We will move from theory to a hands-on, practical workflow. You'll learn how to adapt a pre-trained Whisper checkpoint to a specific domain, significantly improving its accuracy on your target task. Given your experience with PyTorch and the Hugging Face ecosystem, we'll focus on the specific nuances of the audio domain.
Why Fine-Tune Whisper?
Whisper's pre-training on web-scale data gives it a broad understanding of many languages and topics. However, this breadth can sometimes lack depth in niche areas. Fine-tuning allows us to:
- Adapt to a new domain: Teach the model specialized vocabulary, such as medical terms, legal jargon, or brand names.
- Improve performance on specific accents: Increase accuracy for accents or dialects that are underrepresented in the pre-training data.
- Boost accuracy for low-resource languages: Even with its multilingual training, Whisper's performance on some languages can be dramatically improved with just a few hours of high-quality, language-specific data.
The Fine-Tuning Workflow in Hugging Face
We will follow a standard, end-to-end machine learning pipeline using the transformers, datasets, and evaluate libraries. The process can be broken down into three main stages:
- Data Preparation: Loading a custom dataset, processing audio, and tokenizing text.
- Trainer Configuration: Setting up the model, a data collator, evaluation metrics, and training arguments.
- Training and Evaluation: Launching the training job and assessing the fine-tuned model's performance.
Let's begin.
1. Data Preparation
The foundation of successful fine-tuning is a well-prepared dataset. For this lesson, imagine you're adapting Whisper to recognize technical AI/ML terminology it often misspells.
Fine tuning Whisper for Speech Transcription
This video provides an excellent practical demonstration of creating a small, custom dataset to teach Whisper new words. It perfectly illustrates the motivation for fine-tuning.
Watch the segment from 20:13 to 29:45. The presenter walks through the process of creating a custom training and validation set by recording audio and then generating and correcting initial transcripts. This highlights the importance of having accurate audio-text pairs.
As you saw, the process involves:
- Recording Audio: Creating audio files (
.mp3,.wav) containing the target vocabulary. - Creating Transcripts: Generating an initial transcript (even using the base Whisper model) and then meticulously hand-correcting it. The accuracy of this ground truth is critical.
- Structuring the Data: Organizing your data into
audioandsentencecolumns, which is a common format for Hugging Facedatasets. You can load this from local files or, for better collaboration and reproducibility, push it to the Hugging Face Hub.
Loading and Pre-processing
Once you have your audio-text pairs, you'll use the Hugging Face ecosystem to prepare them for the model. This involves loading the feature extractor and tokenizer, which are conveniently wrapped in a single WhisperProcessor.
Fine-tuning the ASR model - Hugging Face Audio Course
The Hugging Face Audio Course provides a clear, code-driven walkthrough of the entire fine-tuning process. We'll use it as our primary guide. This first section covers loading the dataset and preparing the processor.
Please read the sections 'Load Dataset', 'Feature Extractor, Tokenizer and Processor', and 'Pre-Process the Data'. Focus on how the load_dataset function is used, how the WhisperProcessor is instantiated with specific language and task arguments, and the function used to map over the dataset to create input_features and labels.
Let's consolidate the key steps from the reading:
A. Load the Dataset
Using datasets, you load your audio data. A crucial first step is to ensure the audio is at the correct sampling rate. Whisper was pre-trained on audio sampled at 16,000 Hz. Your custom audio might be at a different rate (e.g., 48,000 Hz). You can use the cast_column method to resample it on-the-fly.
from datasets import load_dataset, Audio
# Assuming your custom dataset is on the Hub
# You could also use `load_dataset("audiofolder", data_dir="path/to/your/data")`
dataset = load_dataset("your-username/your-custom-dataset")
# Resample to 16kHz
dataset = dataset.cast_column("audio", Audio(sampling_rate=16000))
B. Instantiate the Processor
The WhisperProcessor combines the feature extractor and tokenizer. You must specify the language and task to ensure the correct special tokens are prepended to the label sequences during tokenization.
from transformers import WhisperProcessor
processor = WhisperProcessor.from_pretrained("openai/whisper-small", language="English", task="transcribe")
C. Create the Preparation Function
You then create a function that will be applied to every example in your dataset. This function:
- Takes a batch of data.
- Uses the feature extractor to convert the 1D audio array into a 2D log-Mel spectrogram (
input_features). - Uses the tokenizer to convert the target text string into a sequence of token IDs (
labels).
Here is a typical implementation:
def prepare_dataset(batch):
# load and resample audio data from 48 to 16kHz
audio = batch["audio"]
# compute log-Mel input features from input audio array
batch["input_features"] = processor.feature_extractor(audio["array"], sampling_rate=audio["sampling_rate"]).input_features[0]
# encode target text to label ids
batch["labels"] = processor.tokenizer(batch["sentence"]).input_ids
return batch
Finally, you apply this function using .map() and optionally filter out any samples longer than 30 seconds, as Whisper processes fixed 30-second chunks.
dataset = dataset.map(prepare_dataset, remove_columns=dataset.column_names["train"])
2. Trainer Configuration
With our data prepared, we now configure the components required by the Seq2SeqTrainer.
Data Collator
This is a subtle but critical component for sequence-to-sequence speech models. The input_features (spectrograms) and labels (token IDs) need to be handled differently when creating a batch.
- Input Features: Already a fixed size (30s spectrogram), so they just need to be stacked into a tensor.
- Labels: Are of variable length. They need to be padded to the length of the longest sequence in the batch. The padding tokens should be replaced with
-100so they are ignored by the cross-entropy loss function.
Fine-Tune Whisper For Multilingual ASR with Transformers
The Hugging Face blog post on fine-tuning Whisper provides an excellent, well-commented implementation of the required data collator.
Read the sub-section 'Define a Data Collator'. Study the DataCollatorSpeechSeq2SeqWithPadding class to understand how it processes input_features and labels separately.
Here's the data collator implementation for your reference:
import torch
from dataclasses import dataclass
from typing import Any, Dict, List, Union
@dataclass
class DataCollatorSpeechSeq2SeqWithPadding:
processor: Any
def __call__(self, features: List[Dict[str, Union[List[int], torch.Tensor]]]) -> Dict[str, torch.Tensor]:
# Split inputs and labels since they have to be of different lengths and need different padding methods
input_features = [{"input_features": feature["input_features"]} for feature in features]
batch = self.processor.feature_extractor.pad(input_features, return_tensors="pt")
label_features = [{"input_ids": feature["labels"]} for feature in features]
labels_batch = self.processor.tokenizer.pad(label_features, return_tensors="pt")
# Replace padding with -100 to ignore in loss
labels = labels_batch["input_ids"].masked_fill(labels_batch.attention_mask.ne(1), -100)
# If bos token is appended in previous tokenization step, cut it here
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)
Evaluation Metric: Word Error Rate (WER)
The standard metric for ASR is Word Error Rate (WER). It measures the number of substitutions, deletions, and insertions required to transform the predicted text into the reference text, divided by the number of words in the reference. A lower WER is better.
The evaluate library makes this easy. We define a compute_metrics function that will be called by the Trainer at each evaluation step.
import evaluate
metric = evaluate.load("wer")
def compute_metrics(pred):
pred_ids = pred.predictions
label_ids = pred.label_ids
# Replace -100 with the pad_token_id
label_ids[label_ids == -100] = processor.tokenizer.pad_token_id
# Decode predictions and labels
pred_str = processor.tokenizer.batch_decode(pred_ids, skip_special_tokens=True)
label_str = processor.tokenizer.batch_decode(label_ids, skip_special_tokens=True)
# Compute WER
wer = 100 * metric.compute(predictions=pred_str, references=label_str)
return {"wer": wer}
Loading the Model and Training Arguments
We load the pre-trained model checkpoint and define our training schedule using Seq2SeqTrainingArguments.
from transformers import WhisperForConditionalGeneration, Seq2SeqTrainingArguments
# Load the model
model = WhisperForConditionalGeneration.from_pretrained("openai/whisper-small")
# Force the model to generate in the correct language
model.config.forced_decoder_ids = processor.get_decoder_prompt_ids(language="english", task="transcribe")
model.config.suppress_tokens = []
# Define training arguments
training_args = Seq2SeqTrainingArguments(
output_dir="./whisper-small-custom", # Directory to save model checkpoints
per_device_train_batch_size=16,
gradient_accumulation_steps=1, # Increase for smaller GPUs
learning_rate=1e-5,
warmup_steps=500,
max_steps=4000,
gradient_checkpointing=True, # Saves VRAM
fp16=True, # Enables mixed-precision training
evaluation_strategy="steps",
per_device_eval_batch_size=8,
predict_with_generate=True,
generation_max_length=225,
save_steps=1000,
eval_steps=1000,
logging_steps=25,
report_to=["tensorboard"],
load_best_model_at_end=True,
metric_for_best_model="wer",
greater_is_better=False,
push_to_hub=True,
)
A Note on Efficiency: LoRA
Fully fine-tuning all 244 million parameters of whisper-small can be computationally expensive. A more efficient method is Parameter-Efficient Fine-Tuning (PEFT), such as Low-Rank Adaptation (LoRA). LoRA freezes the pre-trained model weights and injects trainable rank-decomposition matrices into each layer of the Transformer. This drastically reduces the number of trainable parameters (often >99%), leading to faster training and lower memory usage.
For a practitioner like you, understanding PEFT is crucial. The Trelis Research video provides a good overview of applying LoRA.
Fine tuning Whisper for Speech Transcription
Let's revisit the Trelis Research video to see how LoRA is applied in practice for parameter-efficient fine-tuning.
Watch from 38:20 to 41:30. This segment shows how to define a LoraConfig and use the get_peft_model function from the peft library to wrap the base model. Notice how few parameters are actually being trained.
3. Training and Evaluation
Finally, we bring everything together in the Seq2SeqTrainer and start the training process.
from transformers import Seq2SeqTrainer
trainer = Seq2SeqTrainer(
args=training_args,
model=model,
train_dataset=dataset["train"],
eval_dataset=dataset["test"],
data_collator=data_collator,
compute_metrics=compute_metrics,
tokenizer=processor.feature_extractor, # Pass the processor for generation
)
# Start training!
trainer.train()
# Push the final model to the Hub
trainer.push_to_hub()
During training, you'll monitor the validation loss and, most importantly, the WER. You should see the WER on your validation set decrease significantly as the model adapts to your custom data.
After training is complete, you can load your fine-tuned model using the pipeline for easy inference and see the improved performance on your specialized vocabulary.
from transformers import pipeline
pipe = pipeline(model="your-username/whisper-small-custom")
# Now test it on an audio file with the new terminology
result = pipe("path/to/test_audio.wav")
print(result["text"])
Conclusion
In this lesson, you have walked through the complete, practical pipeline for fine-tuning a Whisper model on a custom dataset using the Hugging Face ecosystem. This process empowers you to transform Whisper from a general-purpose ASR tool into a specialized model tailored to your specific needs.
Key Takeaways:
- Fine-tuning adapts Whisper to new domains, accents, or languages, significantly improving performance where the base model may falter.
- The workflow involves Data Preparation (loading, resampling, processing), Trainer Configuration (defining a data collator, metrics, model, and arguments), and Training.
- The
DataCollatorSpeechSeq2SeqWithPaddingis essential for correctly batching audio features and text labels. - Word Error Rate (WER) is the standard metric for evaluating ASR model performance.
- Techniques like PEFT (LoRA) offer a computationally efficient alternative to full fine-tuning, which is highly valuable in practice.
Preview of the Next Lesson:
Whisper benefits from a powerful internal language model learned during its massive pre-training, an approach known as deep fusion. However, many other powerful ASR architectures, particularly those based on Connectionist Temporal Classification (CTC), do not have this internal LM. In our next lesson, "Explain how a language model can be integrated with a CTC-based acoustic model using shallow fusion," we will explore how to enhance these models by combining their acoustic predictions with an external, separately trained language model. This will broaden your understanding of the different architectural strategies in the ASR landscape.