@@ -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
0 commit comments