File size: 4,152 Bytes
7894cb8 e498f5b c0c0c57 e498f5b 7894cb8 e498f5b f67621d e498f5b c0c0c57 e498f5b 389df6d e498f5b | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | 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)
|