whisper-tiny-en-US / train.py
tomerz14's picture
Rename train_unit5.py to train.py
650aa7f verified
Raw
History Blame Contribute Delete
8.06 kB
"""
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)