Spaces:
Running on Zero
Running on Zero
| 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 | |