Skip to content
This repository was archived by the owner on Apr 8, 2025. It is now read-only.

Commit d4311b3

Browse files
Vectorize Question Answering Prediction Head (#603)
* First attempt vectorize qa ph * Add batch support * add date to component test output * fix test * fix nq processor * Add visualisation of per component benchmark * Vectorize invalid token mask * Remove unused variable * Truncate span_mask * Fix docstring Co-authored-by: Timo Moeller <timo.moeller@deepset.ai> Former-commit-id: 78fb5cf Former-commit-id: 1847e4164ce788dc997fd56f807662a853052249
1 parent 4fffc86 commit d4311b3

6 files changed

Lines changed: 90 additions & 71 deletions

File tree

‎farm/data_handler/input_features.py‎

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -369,7 +369,7 @@ def samples_to_features_bert_lm(sample, max_seq_len, tokenizer, next_sent_pred=T
369369
return [feature_dict]
370370

371371

372-
def sample_to_features_qa(sample, tokenizer, max_seq_len, sp_toks_start, sp_toks_mid,
372+
def sample_to_features_qa(sample, tokenizer, max_seq_len, sp_toks_start, sp_toks_mid, sp_toks_end,
373373
answer_type_list=None, max_answers=6):
374374
""" Prepares data for processing by the model. Supports cases where there are
375375
multiple answers for the one question/document pair. max_answers is by default set to 6 since
@@ -472,6 +472,15 @@ def sample_to_features_qa(sample, tokenizer, max_seq_len, sp_toks_start, sp_toks
472472
# tokens are attended to.
473473
padding_mask = [1] * len(input_ids)
474474

475+
# The passage mask has 1 for tokens that are valid start or ends for QA spans.
476+
# 0s are assigned to question tokens, mid special tokens, end special tokens and padding
477+
# Note that start special tokens are assigned 1 since they can be chosen for a no_answer prediction
478+
span_mask = [1] * sp_toks_start
479+
span_mask += [0] * question_len_t
480+
span_mask += [0] * sp_toks_mid
481+
span_mask += [1] * passage_len_t
482+
span_mask += [0] * sp_toks_end
483+
475484
# Pad up to the sequence length. For certain models, the pad token id is not 0 (e.g. Roberta where it is 1)
476485
pad_idx = tokenizer.pad_token_id
477486
padding = [pad_idx] * (max_seq_len - len(input_ids))
@@ -481,6 +490,7 @@ def sample_to_features_qa(sample, tokenizer, max_seq_len, sp_toks_start, sp_toks
481490
padding_mask += zero_padding
482491
segment_ids += zero_padding
483492
start_of_word += zero_padding
493+
span_mask += zero_padding
484494

485495
# The XLM-Roberta tokenizer generates a segment_ids vector that separates the first sequence from the second.
486496
# However, when this is passed in to the forward fn of the Roberta model, it throws an error since
@@ -500,7 +510,8 @@ def sample_to_features_qa(sample, tokenizer, max_seq_len, sp_toks_start, sp_toks
500510
"start_of_word": start_of_word,
501511
"labels": labels,
502512
"id": sample_id,
503-
"seq_2_start_t": seq_2_start_t}
513+
"seq_2_start_t": seq_2_start_t,
514+
"span_mask": span_mask}
504515
return [feature_dict]
505516

506517

‎farm/data_handler/processor.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1268,6 +1268,7 @@ def _sample_to_features(self, sample) -> dict:
12681268
max_seq_len=self.max_seq_len,
12691269
sp_toks_start=self.sp_toks_start,
12701270
sp_toks_mid=self.sp_toks_mid,
1271+
sp_toks_end=self.sp_toks_end,
12711272
max_answers=self.max_answers)
12721273
return features
12731274

@@ -1572,6 +1573,7 @@ def _sample_to_features(self, sample: Sample) -> dict:
15721573
max_seq_len=self.max_seq_len,
15731574
sp_toks_start=self.sp_toks_start,
15741575
sp_toks_mid=self.sp_toks_mid,
1576+
sp_toks_end=self.sp_toks_end,
15751577
answer_type_list=self.answer_type_list,
15761578
max_answers=self.max_answers)
15771579
return features

‎farm/modeling/prediction_head.py‎

Lines changed: 30 additions & 66 deletions
Original file line numberDiff line numberDiff line change
@@ -1056,7 +1056,8 @@ def logits_to_loss(self, logits, labels, **kwargs):
10561056
per_sample_loss = (start_loss + end_loss) / 2
10571057
return per_sample_loss
10581058

1059-
def logits_to_preds(self, logits, padding_mask, start_of_word, seq_2_start_t, max_answer_length=1000, **kwargs):
1059+
def logits_to_preds(self, logits, span_mask, start_of_word,
1060+
seq_2_start_t, max_answer_length=1000, **kwargs):
10601061
"""
10611062
Get the predicted index of start and end token of the answer. Note that the output is at token level
10621063
and not word level. Note also that these logits correspond to the tokens of a sample
@@ -1077,7 +1078,6 @@ def logits_to_preds(self, logits, padding_mask, start_of_word, seq_2_start_t, ma
10771078
# Calculate a few useful variables
10781079
batch_size = start_logits.size()[0]
10791080
max_seq_len = start_logits.shape[1] # target dim
1080-
n_non_padding = torch.sum(padding_mask, dim=1)
10811081

10821082
# get scores for all combinations of start and end logits => candidate answers
10831083
start_matrix = start_logits.unsqueeze(2).expand(-1, -1, max_seq_len)
@@ -1087,20 +1087,26 @@ def logits_to_preds(self, logits, padding_mask, start_of_word, seq_2_start_t, ma
10871087
# disqualify answers where end < start
10881088
# (set the lower triangular matrix to low value, excluding diagonal)
10891089
indices = torch.tril_indices(max_seq_len, max_seq_len, offset=-1, device=start_end_matrix.device)
1090-
start_end_matrix[:, indices[0][:], indices[1][:]] = -999
1090+
start_end_matrix[:, indices[0][:], indices[1][:]] = -888
10911091

1092-
# disqualify answers where start=0, but end != 0
1093-
start_end_matrix[:, 0, 1:] = -999
1094-
1095-
# TODO continue vectorization of valid_answer_idxs
1096-
# # disqualify where answers < seq_2_start_t and idx != 0
1097-
# # disqualify where answer falls into padding
1098-
# # seq_2_start_t can be different when 2 different questions are handled within one batch
1099-
# # n_non_padding can be different on sample level, too
1100-
# for i in range(batch_size):
1101-
# start_end_matrix[i, 1:seq_2_start_t[i], 1:seq_2_start_t[i]] = -888
1102-
# start_end_matrix[i, n_non_padding[i]-1:, n_non_padding[i]-1:] = -777
1092+
# disqualify answers where answer span is greater than max_answer_length
1093+
# (set the upper triangular matrix to low value, excluding diagonal)
1094+
indices_long_span = torch.triu_indices(max_seq_len, max_seq_len, offset=max_answer_length, device=start_end_matrix.device)
1095+
start_end_matrix[:, indices_long_span[0][:], indices_long_span[1][:]] = -777
11031096

1097+
# disqualify answers where start=0, but end != 0
1098+
start_end_matrix[:, 0, 1:] = -666
1099+
1100+
# Turn 1d span_mask vectors into 2d span_mask along 2 different axes
1101+
# span mask has:
1102+
# 0 for every position that is never a valid start or end index (question tokens, mid and end special tokens, padding)
1103+
# 1 everywhere else
1104+
span_mask_start = span_mask.unsqueeze(2).expand(-1, -1, max_seq_len)
1105+
span_mask_end = span_mask.unsqueeze(1).expand(-1, max_seq_len, -1)
1106+
span_mask_2d = span_mask_start + span_mask_end
1107+
# disqualify spans where either start or end is on an invalid token
1108+
invalid_indices = torch.nonzero((span_mask_2d != 2), as_tuple=True)
1109+
start_end_matrix[invalid_indices[0][:], invalid_indices[1][:], invalid_indices[2][:]] = -999
11041110

11051111
# Sort the candidate answers by their score. Sorting happens on the flattened matrix.
11061112
# flat_sorted_indices.shape: (batch_size, max_seq_len^2, 1)
@@ -1114,20 +1120,16 @@ def logits_to_preds(self, logits, padding_mask, start_of_word, seq_2_start_t, ma
11141120
end_indices = flat_sorted_indices % max_seq_len
11151121
sorted_candidates = torch.cat((start_indices, end_indices), dim=2)
11161122

1117-
# Get the n_best candidate answers for each sample that are valid (via some heuristic checks)
1123+
# Get the n_best candidate answers for each sample
11181124
for sample_idx in range(batch_size):
11191125
sample_top_n = self.get_top_candidates(sorted_candidates[sample_idx],
11201126
start_end_matrix[sample_idx],
1121-
n_non_padding[sample_idx].item(),
1122-
max_answer_length,
1123-
seq_2_start_t[sample_idx].item(),
11241127
sample_idx)
11251128
all_top_n.append(sample_top_n)
11261129

11271130
return all_top_n
11281131

1129-
def get_top_candidates(self, sorted_candidates, start_end_matrix,
1130-
n_non_padding, max_answer_length, seq_2_start_t, sample_idx):
1132+
def get_top_candidates(self, sorted_candidates, start_end_matrix, sample_idx):
11311133
""" Returns top candidate answers as a list of Span objects. Operates on a matrix of summed start and end logits.
11321134
This matrix corresponds to a single sample (includes special tokens, question tokens, passage tokens).
11331135
This method always returns a list of len n_best + 1 (it is comprised of the n_best positive answers along with the one no_answer)"""
@@ -1147,16 +1149,14 @@ def get_top_candidates(self, sorted_candidates, start_end_matrix,
11471149
# Ignore no_answer scores which will be extracted later in this method
11481150
if start_idx == 0 and end_idx == 0:
11491151
continue
1150-
# Check that the candidate's indices are valid and save them if they are
1151-
if self.valid_answer_idxs(start_idx, end_idx, n_non_padding, max_answer_length, seq_2_start_t):
1152-
score = start_end_matrix[start_idx, end_idx].item()
1153-
top_candidates.append(QACandidate(offset_answer_start=start_idx,
1154-
offset_answer_end=end_idx,
1155-
score=score,
1156-
answer_type="span",
1157-
offset_unit="token",
1158-
aggregation_level="passage",
1159-
passage_id=sample_idx))
1152+
score = start_end_matrix[start_idx, end_idx].item()
1153+
top_candidates.append(QACandidate(offset_answer_start=start_idx,
1154+
offset_answer_end=end_idx,
1155+
score=score,
1156+
answer_type="span",
1157+
offset_unit="token",
1158+
aggregation_level="passage",
1159+
passage_id=sample_idx))
11601160

11611161
no_answer_score = start_end_matrix[0, 0].item()
11621162
top_candidates.append(QACandidate(offset_answer_start=0,
@@ -1169,42 +1169,6 @@ def get_top_candidates(self, sorted_candidates, start_end_matrix,
11691169

11701170
return top_candidates
11711171

1172-
@staticmethod
1173-
def valid_answer_idxs(start_idx, end_idx, n_non_padding, max_answer_length, seq_2_start_t):
1174-
""" Returns True if the supplied index span is a valid prediction. The indices being provided
1175-
should be on sample/passage level (special tokens + question_tokens + passag_tokens)
1176-
and not document level"""
1177-
1178-
# This function can seriously slow down inferencing and eval. In the future this function will be completely vectorized
1179-
# Continue if start or end label points to a padding token
1180-
if start_idx < seq_2_start_t and start_idx != 0:
1181-
return False
1182-
if end_idx < seq_2_start_t and end_idx != 0:
1183-
return False
1184-
# The -1 is to stop the idx falling on a final special token
1185-
# TODO: this makes the assumption that there is a special token that comes at the end of the second sequence
1186-
if start_idx >= n_non_padding - 1:
1187-
return False
1188-
if end_idx >= n_non_padding - 1:
1189-
return False
1190-
1191-
# # Check if start comes after end
1192-
# # Handled on matrix level by: start_end_matrix[:, indices[0][1:], indices[1][1:]] = -999
1193-
# if end_idx < start_idx:
1194-
# return False
1195-
1196-
# # If one of the two indices is 0, the other must also be 0
1197-
# # Handled on matrix level by setting: start_end_matrix[:, 0, 1:] = -999
1198-
# if start_idx == 0 and end_idx != 0:
1199-
# return False
1200-
# if start_idx != 0 and end_idx == 0:
1201-
# return False
1202-
1203-
length = end_idx - start_idx + 1
1204-
if length > max_answer_length:
1205-
return False
1206-
return True
1207-
12081172
def formatted_preds(self, logits=None, preds=None, baskets=None, **kwargs):
12091173
""" Takes a list of passage level predictions, each corresponding to one sample, and converts them into document level
12101174
predictions. Leverages information in the SampleBaskets. Assumes that we are being passed predictions from
Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,39 @@
1+
2+
<html>
3+
<head>
4+
<script type="text/javascript" src="https://www.gstatic.com/charts/loader.js"></script>
5+
<script type="text/javascript">
6+
google.charts.load('current', {'packages':['bar']});
7+
google.charts.setOnLoadCallback(drawChart);
8+
9+
function drawChart() {
10+
var data = google.visualization.arrayToDataTable(
11+
[
12+
["Name", "preproc","language_model","prediction_head"],
13+
['deepset/minilm-uncased-squad2', 12.277034912109375, 5.79623876953125, 1.5562604980468748], ['deepset/roberta-base-squad2', 12.380782958984376, 13.71148828125, 1.5372104492187502], ['deepset/bert-base-cased-squad2', 9.938722900390625, 15.864041992187499, 1.6085009765625005], ['deepset/bert-large-uncased-whole-word-masking-squad2', 9.692403808593749, 45.28969921875, 1.785435546875], ['deepset/xlm-roberta-large-squad2', 8.079997680664063, 48.489154296875, 1.974138671875]
14+
]);
15+
16+
var options = {
17+
chart: {
18+
title: 'QA Model Speed Comparison',
19+
subtitle: 'Time per Component',
20+
},
21+
bars: 'horizontal', // Required for Material Bar Charts.
22+
isStacked: true,
23+
height: 300,
24+
legend: {position: 'top', maxLines: 3},
25+
hAxis: {minValue: 0}
26+
27+
};
28+
29+
var chart = new google.charts.Bar(document.getElementById('barchart_material'));
30+
31+
chart.draw(data, google.charts.Bar.convertOptions(options));
32+
}
33+
</script>
34+
</head>
35+
<body>
36+
<div id="barchart_material" style="width: 900px; height: 500px;"></div>
37+
</body>
38+
</html>
39+

‎test/benchmarks/question_answering_components.py‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -10,16 +10,18 @@
1010
from tqdm import tqdm
1111
import logging
1212
import json
13+
from datetime import date
1314

14-
logger = logging.getLogger(__name__)
1515

16+
logger = logging.getLogger(__name__)
1617

1718
task_type = "question_answering"
1819
sample_file = "samples/question_answering_sample.txt"
1920
questions_file = "samples/question_answering_questions.txt"
2021
num_processes = 1
2122
passages_per_char = 2400 / 1000000 # numerator is number of passages when 1mill chars paired with one of the questions, msl 384, doc stride 128
22-
output_file = "results_component_test_24_09_20.csv"
23+
date_str = date.today().strftime("%d_%m_%Y")
24+
output_file = f"results_component_test_{date_str}.csv"
2325

2426
params = {
2527
"modelname": ["deepset/bert-base-cased-squad2", "deepset/minilm-uncased-squad2", "deepset/roberta-base-squad2", "deepset/bert-large-uncased-whole-word-masking-squad2", "deepset/xlm-roberta-large-squad2"],

‎test/test_input_features.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@
99
MODEL = "roberta-base"
1010
SP_TOKENS_START = 1
1111
SP_TOKENS_MID = 2
12+
SP_TOKENS_END = 1
1213

1314
def to_list(x):
1415
try:
@@ -32,7 +33,7 @@ def test_sample_to_features_qa(caplog):
3233
curr_id = "-".join([str(x) for x in features_gold["id"]])
3334

3435
s = Sample(id=curr_id, clear_text=clear_text, tokenized=tokenized)
35-
features = sample_to_features_qa(s, tokenizer, max_seq_len, SP_TOKENS_START, SP_TOKENS_MID)[0]
36+
features = sample_to_features_qa(s, tokenizer, max_seq_len, SP_TOKENS_START, SP_TOKENS_MID, SP_TOKENS_END)[0]
3637
features = to_list(features)
3738

3839
keys = features_gold.keys()

0 commit comments

Comments
 (0)