From ac99a47a45bc80b424f48ce9e50de59b209af992 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 21 Jan 2024 18:52:15 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 747ef0793a..ff6b0afdda 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -189,7 +189,7 @@ def LlamaAttention_fast_forward( bsz, q_len, _ = hidden_states.size() # Check for inference - if use_cache and past_key_value is not None and q_len == 1: + if past_key_value is not None and q_len == 1: A, past_key_value = LlamaAttention_fast_forward_inference( self, hidden_states, @@ -305,8 +305,7 @@ def LlamaDecoderLayer_fast_forward( """ bsz, q_len, hd = hidden_states.size() - if (not self.training and q_len == 1): - print(1) + if (past_key_value is not None and q_len == 1): # Self Attention residual = hidden_states hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states) @@ -525,7 +524,7 @@ def LlamaModel_fast_forward( pass bsz, q_len, hd = hidden_states.size() - if (not self.training and q_len == 1): + if (past_key_value is not None and q_len == 1): hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states) else: hidden_states = fast_rms_layernorm(self.norm, hidden_states)