namnh113/Question_Answering
0
1import torch2import torch.nn as nn3from transformers.models.roberta.modeling_roberta import *4 5 6class MRCQuestionAnswering(RobertaPreTrainedModel):7 config_class = RobertaConfig8 9 def _reorder_cache(self, past, beam_idx):10 pass11 12 _keys_to_ignore_on_load_unexpected = [r"pooler"]13 _keys_to_ignore_on_load_missing = [r"position_ids"]14 15 def __init__(self, config):16 super().__init__(config)17 self.num_labels = config.num_labels18 19 self.roberta = RobertaModel(config, add_pooling_layer=False)20 self.qa_outputs = nn.Linear(config.hidden_size, config.num_labels)21 22 self.init_weights()23 24 def forward(25 self,26 input_ids=None,27 words_lengths=None,28 start_idx=None,29 end_idx=None,30 attention_mask=None,31 token_type_ids=None,32 position_ids=None,33 head_mask=None,34 inputs_embeds=None,35 start_positions=None,36 end_positions=None,37 span_answer_ids=None,38 output_attentions=None,39 output_hidden_states=None,40 return_dict=None,41 ):42 return_dict = return_dict if return_dict is not None else self.config.use_return_dict43 44 outputs = self.roberta(45 input_ids,46 attention_mask=attention_mask,47 token_type_ids=token_type_ids,48 position_ids=position_ids,49 head_mask=head_mask,50 inputs_embeds=inputs_embeds,51 output_attentions=output_attentions,52 output_hidden_states=output_hidden_states,53 return_dict=return_dict,54 )55 56 sequence_output = outputs[0]57 58 context_embedding = sequence_output59 60 # Compute align word sub_word matrix61 batch_size = input_ids.shape[0]62 max_sub_word = input_ids.shape[1]63 max_word = words_lengths.shape[1]64 align_matrix = torch.zeros((batch_size, max_word, max_sub_word))65 66 for i, sample_length in enumerate(words_lengths):67 for j in range(len(sample_length)):68 start_idx = torch.sum(sample_length[:j])69 align_matrix[i][j][start_idx: start_idx + sample_length[j]] = 1 if sample_length[j] > 0 else 070 71 align_matrix = align_matrix.to(context_embedding.device)72 # Combine sub_word features to make word feature73 context_embedding_align = torch.bmm(align_matrix, context_embedding)74 75 logits = self.qa_outputs(context_embedding_align)76 start_logits, end_logits = logits.split(1, dim=-1)77 start_logits = start_logits.squeeze(-1).contiguous()78 end_logits = end_logits.squeeze(-1).contiguous()79 80 total_loss = None81 if start_positions is not None and end_positions is not None:82 # If we are on multi-GPU, split add a dimension83 if len(start_positions.size()) > 1:84 start_positions = start_positions.squeeze(-1)85 if len(end_positions.size()) > 1:86 end_positions = end_positions.squeeze(-1)87 # sometimes the start/end positions are outside our model inputs, we ignore these terms88 ignored_index = start_logits.size(1)89 start_positions = start_positions.clamp(0, ignored_index)90 end_positions = end_positions.clamp(0, ignored_index)91 92 loss_fct = nn.CrossEntropyLoss(ignore_index=ignored_index)93 start_loss = loss_fct(start_logits, start_positions)94 end_loss = loss_fct(end_logits, end_positions)95 total_loss = (start_loss + end_loss) / 296 97 if not return_dict:98 output = (start_logits, end_logits) + outputs[2:]99 return ((total_loss,) + output) if total_loss is not None else output100 101 return QuestionAnsweringModelOutput(102 loss=total_loss,103 start_logits=start_logits,104 end_logits=end_logits,105 hidden_states=outputs.hidden_states,106 attentions=outputs.attentions,107 )