File size: 1,759 Bytes
c3197e2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from dataclasses import dataclass

import torch
import torch.nn as nn
from transformers import HubertConfig, HubertModel, PreTrainedModel
from transformers.utils import ModelOutput

from .configuration_mhubert_ipa_ctc_ft import MHuBERTIPACTCFTConfig


@dataclass
class MHuBERTIPACTCFTOutput(ModelOutput):
    logits: torch.Tensor = None
    hidden_states: torch.Tensor = None


class MHuBERTIPACTCFTModel(PreTrainedModel):
    config_class = MHuBERTIPACTCFTConfig
    base_model_prefix = "mhubert_ipa_ctc_ft"

    def __init__(self, config):
        super().__init__(config)
        arch = config.architecture
        backbone_cfg = HubertConfig.from_dict(config.backbone_config)
        self.blank_id = int(arch["blank_id"])
        self.backbone = HubertModel(backbone_cfg)
        self.proj = nn.Linear(arch["input_dim"], arch["proj_dim"])
        self.lstm = nn.LSTM(
            arch["proj_dim"],
            arch["lstm_hidden"],
            num_layers=arch["lstm_layers"],
            bidirectional=arch["lstm_bidirectional"],
            batch_first=True,
            dropout=arch["dropout"] if arch["lstm_layers"] > 1 else 0.0,
        )
        self.drop = nn.Dropout(arch["dropout"])
        out_dim = arch["lstm_hidden"] * (2 if arch["lstm_bidirectional"] else 1)
        self.head = nn.Linear(out_dim, arch["output_dim"])
        self.post_init()

    def forward(self, input_values, attention_mask=None, **kwargs):
        backbone_out = self.backbone(input_values=input_values, attention_mask=attention_mask, **kwargs)
        x = self.proj(backbone_out.last_hidden_state)
        out, _ = self.lstm(x)
        logits = self.head(self.drop(out))
        return MHuBERTIPACTCFTOutput(logits=logits, hidden_states=backbone_out.last_hidden_state)