bilstm-sequence-labeler / modeling_bilstm.py
ILoveBacteria's picture
Fix: Invoke self.post_init() to satisfy all_tied_weights_keys requirements
9dba0cb verified
Raw
History Blame Contribute Delete
1.67 kB
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)
}