JEDI / modeling.py
nicpopovic's picture
commit
595f017
Raw
History Blame Contribute Delete
7.65 kB
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