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