import librosa import torch from dataclasses import dataclass from typing import Any from datasets import load_dataset from transformers import ( AutoTokenizer, Trainer, TrainingArguments, Wav2Vec2ForCTC, Wav2Vec2Processor, ) MODEL_ID = "facebook/wav2vec2-base" DATASET_ID = "keithito/lj_speech" @dataclass class DataCollatorCTCWithPadding: processor: Wav2Vec2Processor padding: bool | str = True def __call__(self, features: list[dict[str, Any]]) -> dict[str, torch.Tensor]: input_features = [{"input_values": f["input_values"]} for f in features] label_features = [{"input_ids": f["labels"]} for f in features] batch = self.processor.feature_extractor.pad( input_features, padding=self.padding, return_tensors="pt" ) labels_batch = self.processor.tokenizer.pad( label_features, padding=self.padding, return_tensors="pt" ) labels = labels_batch["input_ids"].masked_fill( labels_batch.attention_mask.ne(1), -100 ) batch["labels"] = labels return batch def prepare_dataset(batch, processor): audio = batch["audio"] target_sr = processor.feature_extractor.sampling_rate array = audio["array"] if audio["sampling_rate"] != target_sr: array = librosa.resample( array, orig_sr=audio["sampling_rate"], target_sr=target_sr, res_type="kaiser_fast", ) batch["input_values"] = processor(array, sampling_rate=target_sr).input_values[0] batch["input_length"] = len(batch["input_values"]) batch["labels"] = processor(text=batch["normalized_text"]).input_ids return batch def compute_metrics(pred): from jiwer import wer pred_logits = pred.predictions pred_ids = torch.argmax(torch.tensor(pred_logits), dim=-1) pred_str = [s if len(s) > 0 else "-" for s in tokenizer.batch_decode(pred_ids)] label_ids = pred.label_ids.copy() label_ids[label_ids == -100] = tokenizer.pad_token_id label_str = tokenizer.batch_decode(label_ids, group_tokens=False) # jiwer.wer(reference, hypothesis) — positional only error = wer(label_str, pred_str) return {"wer": error} if __name__ == "__main__": tokenizer = AutoTokenizer.from_pretrained(MODEL_ID) processor = Wav2Vec2Processor.from_pretrained(MODEL_ID) dataset = load_dataset(DATASET_ID, split="train") dataset = dataset.train_test_split(test_size=0.05, seed=42) dataset = dataset.map( lambda batch: prepare_dataset(batch, processor), remove_columns=dataset.column_names["train"], num_proc=4, ) dataset = dataset.filter( lambda x: x < processor.tokenizer.model_max_length, input_columns=["input_length"], num_proc=4, ) data_collator = DataCollatorCTCWithPadding(processor=processor) model = Wav2Vec2ForCTC.from_pretrained( MODEL_ID, attention_dropout=0.0, hidden_dropout=0.0, feat_proj_dropout=0.0, mask_time_prob=0.0, layerdrop=0.0, ctc_loss_reduction="mean", pad_token_id=processor.tokenizer.pad_token_id, vocab_size=len(processor.tokenizer), ) model.freeze_feature_encoder() training_args = TrainingArguments( output_dir="wav2vec2-ljspeech", per_device_train_batch_size=8, gradient_accumulation_steps=2, learning_rate=3e-4, num_train_epochs=10, warmup_steps=500, fp16=True, save_steps=500, eval_strategy="steps", eval_steps=500, logging_steps=100, save_total_limit=3, report_to=["tensorboard"], gradient_checkpointing=True, ) trainer = Trainer( model=model, args=training_args, train_dataset=dataset["train"], eval_dataset=dataset["test"], data_collator=data_collator, compute_metrics=compute_metrics, processing_class=processor.tokenizer, ) trainer.train() trainer.save_model(training_args.output_dir) processor.save_pretrained(training_args.output_dir)