""" BiLSTM + CRF (+optional CharCNN) for Spanish NER. Install: pip install pytorch-crf """ import torch import torch.nn as nn from transformers import PreTrainedModel from transformers.modeling_outputs import TokenClassifierOutput from .configuration_bilstm import BiLSTMConfig class BiLSTMForTokenClassification(PreTrainedModel): config_class = BiLSTMConfig base_model_prefix = "bilstm" def __init__(self, config): super().__init__(config) self.emb = nn.Embedding(config.vocab_size, config.embed_dim, padding_idx=config.pad_token_id) lstm_in = config.embed_dim if config.use_char_cnn: self.char_emb = nn.Embedding(config.char_vocab_size, config.char_emb_dim, padding_idx=0) self.char_cnn = nn.Conv1d(config.char_emb_dim, config.char_cnn_filters, kernel_size=config.char_cnn_kernel, padding=config.char_cnn_kernel // 2) lstm_in += config.char_cnn_filters self.lstm = nn.LSTM(lstm_in, config.hidden_dim, batch_first=True, bidirectional=True) self.drop = nn.Dropout(config.dropout) self.fc = nn.Linear(config.hidden_dim * 2, config.num_labels) from torchcrf import CRF self.crf = CRF(config.num_labels, batch_first=True) self.post_init() def forward(self, input_ids, char_input_ids=None, attention_mask=None, labels=None, **kwargs): we = self.emb(input_ids) if self.config.use_char_cnn: B, L, W = char_input_ids.shape ce = self.char_emb(char_input_ids.view(B*L, W)).transpose(1, 2) cf, _ = torch.relu(self.char_cnn(ce)).max(dim=2) cf = cf.view(B, L, -1) x = torch.cat([we, cf], dim=-1) else: x = we h, _ = self.lstm(x) h = self.drop(h) emis = self.fc(h) if attention_mask is None: attention_mask = torch.ones(input_ids.shape, dtype=torch.bool, device=input_ids.device) else: attention_mask = attention_mask.bool() loss = None if labels is not None: sl = labels.clone(); sl[sl == -100] = 0 loss = -self.crf(emis, sl, mask=attention_mask, reduction="mean") decoded = self.crf.decode(emis, mask=attention_mask) B, L, K = emis.shape logits = torch.full((B, L, K), -1e4, device=emis.device) for i, seq in enumerate(decoded): for j, t in enumerate(seq): logits[i, j, t] = 0.0 return TokenClassifierOutput(loss=loss, logits=logits) @torch.no_grad() def predict(self, tokens, vocab): """Helper de inferencia. tokens: List[str] o List[List[str]].""" self.eval() device = next(self.parameters()).device if tokens and isinstance(tokens[0], str): tokens = [tokens] word2idx = vocab["word2idx"] char2idx = vocab.get("char2idx", {}) id2tag = {int(k): v for k, v in vocab["id2tag"].items()} unk = self.config.unk_token_id max_word = self.config.max_word_len L = max(len(s) for s in tokens) tok_ids, char_ids, masks = [], [], [] for sent in tokens: ti = [word2idx.get(w, unk) for w in sent] + [0]*(L-len(sent)) ci = [[char2idx.get(c, 1) for c in w[:max_word]] + [0]*(max_word-len(w[:max_word])) for w in sent] ci += [[0]*max_word]*(L-len(sent)) mk = [1]*len(sent) + [0]*(L-len(sent)) tok_ids.append(ti); char_ids.append(ci); masks.append(mk) tok = torch.tensor(tok_ids, dtype=torch.long, device=device) ch = torch.tensor(char_ids, dtype=torch.long, device=device) mk = torch.tensor(masks, dtype=torch.bool, device=device) out = self.forward(input_ids=tok, char_input_ids=ch if self.config.use_char_cnn else None, attention_mask=mk) preds = out.logits.argmax(-1).tolist() result = [] for i, sent in enumerate(tokens): tags = [id2tag.get(int(preds[i][j]), "O") for j in range(len(sent))] result.append(tags) return result