| 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) |
|
|
| |
| 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) |
|
|