CoolFace
Modelpublic

entropy/roberta_zinc_decoder

sourceHugging Faceupdated 1y agoView on Hugging Face
1likes69downloads
conditional_gpt2_model.py43 linesDownload Raw Back to root
1from transformers import GPT2LMHeadModel2 3class ConditionalGPT2LMHeadModel(GPT2LMHeadModel):        4    def prepare_inputs_for_generation(self, input_ids, past_key_values=None, inputs_embeds=None, **kwargs):5        # this is the same as `GPT2LMHeadModel` except `encoder_hidden_states` are added to inputs6        7        token_type_ids = kwargs.get("token_type_ids", None)8        # only last token for inputs_ids if past is defined in kwargs9        if past_key_values:10            input_ids = input_ids[:, -1].unsqueeze(-1)11            if token_type_ids is not None:12                token_type_ids = token_type_ids[:, -1].unsqueeze(-1)13 14        attention_mask = kwargs.get("attention_mask", None)15        position_ids = kwargs.get("position_ids", None)16 17        if attention_mask is not None and position_ids is None:18            # create position_ids on the fly for batch generation19            position_ids = attention_mask.long().cumsum(-1) - 120            position_ids.masked_fill_(attention_mask == 0, 1)21            if past_key_values:22                position_ids = position_ids[:, -1].unsqueeze(-1)23        else:24            position_ids = None25 26        # if `inputs_embeds` are passed, we only want to use them in the 1st generation step27        if inputs_embeds is not None and past_key_values is None:28            model_inputs = {"inputs_embeds": inputs_embeds}29        else:30            model_inputs = {"input_ids": input_ids}31            32        model_inputs['encoder_hidden_states'] = kwargs.get('encoder_hidden_states', None)33 34        model_inputs.update(35            {36                "past_key_values": past_key_values,37                "use_cache": kwargs.get("use_cache"),38                "position_ids": position_ids,39                "attention_mask": attention_mask,40                "token_type_ids": token_type_ids,41            }42        )43        return model_inputs