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