File size: 1,669 Bytes
ebe9f23 9dba0cb ebe9f23 | 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 | import torch
import torch.nn as nn
from transformers import PretrainedConfig, PreTrainedModel
class BiLSTMSequenceLabelerConfig(PretrainedConfig):
model_type = "bilstm_sequence_labeler"
def __init__(self, vocab_size=30000, embedding_dim=300, hidden_dim=128, num_ner_classes=9, num_pos_classes=45, num_chunk_classes=23, **kwargs):
super().__init__(**kwargs)
self.vocab_size = vocab_size
self.embedding_dim = embedding_dim
self.hidden_dim = hidden_dim
self.num_ner_classes = num_ner_classes
self.num_pos_classes = num_pos_classes
self.num_chunk_classes = num_chunk_classes
class HFBiLSTMSequenceLabeler(PreTrainedModel):
config_class = BiLSTMSequenceLabelerConfig
def __init__(self, config):
super().__init__(config)
self.embedding = nn.Embedding(config.vocab_size, config.embedding_dim, padding_idx=0)
self.rnn = nn.LSTM(config.embedding_dim, config.hidden_dim, batch_first=True, bidirectional=True)
rnn_out_dim = config.hidden_dim * 2
self.ner_head = nn.Linear(rnn_out_dim, config.num_ner_classes)
self.pos_head = nn.Linear(rnn_out_dim, config.num_pos_classes)
self.chunk_head = nn.Linear(rnn_out_dim, config.num_chunk_classes)
# CRUCIAL: Hook for newer transformers compatibility (initializes all_tied_weights_keys)
self.post_init()
def forward(self, input_ids, **kwargs):
embedded = self.embedding(input_ids)
rnn_out, _ = self.rnn(embedded)
return {
"ner": self.ner_head(rnn_out),
"pos": self.pos_head(rnn_out),
"chunk": self.chunk_head(rnn_out)
}
|