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)