entropy/roberta_zinc_decoder
169
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