From 82eea75cde8bf4288ebee9fd014b32b612943fff Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 2 Feb 2024 20:29:35 +1100 Subject: [PATCH] torch compile --- unsloth/models/llama.py | 1 - unsloth/models/mistral.py | 34 ++++++++++++++++++++++------------ 2 files changed, 22 insertions(+), 13 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index efcb86012d..fa3c8d1c5a 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -752,7 +752,6 @@ def LlamaForCausalLM_fast_forward( if past_key_value is not None and \ hasattr(self.model.layers[0].self_attn, "paged_attention"): - print(1) outputs = LlamaModel_fast_forward_inference( self.model, input_ids, diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index 7b8302629c..99dc3c1fd9 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -200,18 +200,28 @@ def MistralForCausalLM_fast_forward( # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) self.model._has_no_labels = labels is None - outputs = self.model( - input_ids=input_ids, - causal_mask=causal_mask, - attention_mask=attention_mask, - position_ids=position_ids, - past_key_values=past_key_values, - inputs_embeds=inputs_embeds, - use_cache=use_cache, - output_attentions=output_attentions, - output_hidden_states=output_hidden_states, - return_dict=return_dict, - ) + + if past_key_value is not None and \ + hasattr(self.model.layers[0].self_attn, "paged_attention"): + outputs = LlamaModel_fast_forward_inference( + self.model, + input_ids, + past_key_values, + ) + else: + outputs = self.model( + input_ids=input_ids, + causal_mask=causal_mask, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + inputs_embeds=inputs_embeds, + use_cache=use_cache, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + pass hidden_states = outputs[0] bsz, q_len, hd = hidden_states.shape