CoolFace
Apppublic

namnh113/Question_Answering

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
model.py107 linesDownload Raw Back to root
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        )