import torch import torch.nn as nn import torch.nn.functional as F import random import numpy as np def spans_overlap(span1, span2): # Returns True if span1 and span2 overlap return not (span1[1] <= span2[0] or span1[0] >= span2[1]) def span_contains(span_a, span_b): """Check if span_a contains span_b. Args: span_a (tuple or list): The outer span (start, end). span_b (tuple or list): The inner span to check for containment. Returns: bool: True if span_a contains span_b, False otherwise. """ a_start, a_end = span_a b_start, b_end = span_b return a_start <= b_start and a_end >= b_end def update_salient_spans(salient_spans, spans): # Make a copy to avoid modifying while iterating new_salient_spans = salient_spans[:] for span in spans: for salient in salient_spans: if span_contains(span, salient) and (salient[1]-salient[0])/(span[1]-span[0]) >= 0.9: if span not in new_salient_spans: new_salient_spans.append(span) break return new_salient_spans class IENLIModel(nn.Module): """Wrapper for transformers model.""" def __init__(self, model, tokenizer): super().__init__() self.tokenizer = tokenizer self.encoder = model self.septoken = ("[SEP]", self.tokenizer("[SEP]", add_special_tokens=False)['input_ids'][0]) self.span_reduce = nn.Linear(2*self.encoder.config.hidden_size, self.encoder.config.hidden_size) self.classification_reduce = nn.Linear(2*self.encoder.config.hidden_size, self.encoder.config.hidden_size) self.emb_size = self.encoder.config.hidden_size self.block_size = 64 self.bilinear = nn.Linear(self.encoder.config.hidden_size * 64, 3) self.classification = nn.Linear(2*self.encoder.config.hidden_size, 3) self.is_left = nn.Linear(self.encoder.config.hidden_size, 1) self.is_span = nn.Linear(self.encoder.config.hidden_size, 1) def forward(self, data, k_spans=10): """Predict the label for the input data using the model.""" input_ids = data["input_ids"].to(self.encoder.device) attention_mask = data["attention_mask"].to(self.encoder.device) sep_locations = data["sep_locations"] offset_mapping = data["offset_mapping"] outputs = self.encoder(input_ids=input_ids, attention_mask=attention_mask, output_hidden_states=True) hidden_states = outputs.hidden_states[-1] b1 = torch.tanh(outputs.hidden_states[-1][[i for i in range(len(sep_locations))], sep_locations, :]).view(-1, self.emb_size // self.block_size, self.block_size) b2 = outputs.hidden_states[-1][:, 0, :].view(-1, self.emb_size // self.block_size, self.block_size) bl = (b1.unsqueeze(3) * b2.unsqueeze(2)).view(-1, self.emb_size * self.block_size) prediction_neutral_or_not = self.bilinear(bl) predictions_left = torch.sigmoid(self.is_left(hidden_states)) predictions_positive = torch.where(predictions_left > 0.5) left_topk_indices = [[] for _ in range(predictions_left.shape[0])] for x, y in zip(predictions_positive[0], predictions_positive[1]): if y.item() != 0 and y.item() < sep_locations[x.item()]: left_topk_indices[x.item()].append(y) spans_predicted = [] spans_predicted_entailed = [] spans_predicted_contradicted = [] softmaxes = [] selected_spans_probabilities = [] verdicts = [] for i in range(len(sep_locations)): left_indices = left_topk_indices[i] if len(left_indices) == 0 and not self.training: spans_predicted.append([]) spans_predicted_entailed.append([]) spans_predicted_contradicted.append([]) softmaxes.append([]) selected_spans_probabilities.append([]) verdicts.append(1) continue index = [] left_embeddings = [] right_embeddings = [] lengths = [] for j, left in enumerate(left_indices): right_embeddings.append(hidden_states[i, left_indices[j]:sep_locations[i]+1]) left_embeddings.append(hidden_states[i, left_indices[j]].unsqueeze(0).repeat(len(right_embeddings[-1]), 1)) lengths.append(len(right_embeddings[-1])) index.extend([(left.item(), j) for j in range(left_indices[j], sep_locations[i]+1)]) right_embeddings = torch.cat(right_embeddings) left_embeddings = torch.cat(left_embeddings) span_embeddings_unreduced = torch.cat((left_embeddings, right_embeddings), dim=-1) span_embeddings = self.span_reduce(span_embeddings_unreduced) predictions_span = torch.sigmoid(self.is_span(span_embeddings)) predictions_positive = torch.where(predictions_span > 0.5) spans_topk_indices = [x.item() for x in predictions_positive[0]] spans_top_k = [] span_embeddings_top_k = [] for ind in spans_topk_indices: spans_top_k.append(index[ind]) span_embeddings_top_k.append(span_embeddings_unreduced[ind]) spans_predicted.append(spans_top_k) selected_spans_probabilities.append(predictions_span[spans_topk_indices]) if len(spans_top_k) == 0: spans_predicted_entailed.append([]) spans_predicted_contradicted.append([]) softmaxes.append([]) verdicts.append(1) continue span_embeddings_top_k = torch.stack(span_embeddings_top_k) span_embeddings_class = torch.tanh(self.classification_reduce(span_embeddings_top_k)) b1 = span_embeddings_class.view(-1, self.emb_size // self.block_size, self.block_size) b2 = outputs.hidden_states[-1][i, 0, :].unsqueeze(0).repeat(span_embeddings_top_k.shape[0], 1).view(-1, self.emb_size // self.block_size, self.block_size) bl = (b1.unsqueeze(3) * b2.unsqueeze(2)).view(-1, self.emb_size * self.block_size) prediction = self.bilinear(bl) predictions = torch.argmax(prediction, dim=-1) softmaxes.append([torch.softmax(prediction, dim=-1)[i, x] for i, x in enumerate(predictions)]) spans_predicted_entailed.append(predictions) spans_predicted_contradicted.append([]) if torch.argmax(prediction_neutral_or_not[i]) == 2 and torch.any(predictions == 2): verdicts.append(2) elif torch.argmax(prediction_neutral_or_not[i]) == 0 and torch.any(predictions == 0): verdicts.append(0) else: verdicts.append(1) spans_predicted_char = [] for i, span_list in enumerate(spans_predicted): char_spans = [] for token_start_idx, token_end_idx in span_list: token_start_idx = min(token_start_idx, len(offset_mapping[i]) - 1) token_end_idx = min(token_end_idx, len(offset_mapping[i]) - 1) start_char = offset_mapping[i][token_start_idx][0] end_char = offset_mapping[i][token_end_idx][1] char_spans.append((start_char, end_char)) spans_predicted_char.append(char_spans) return verdicts, (spans_predicted_char, spans_predicted_entailed, spans_predicted_contradicted), softmaxes, selected_spans_probabilities