Isa0's picture
fix arguments in wer() functions
f67621d
Raw
History Blame Contribute Delete
4.15 kB
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)