diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 57f8a42aec..d2251627c5 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -358,7 +358,7 @@ def LlamaDecoderLayer_fast_forward( (see `past_key_values`). past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states """ - if past_key_value is not None: + if False:#past_key_value is not None: do_prefill = not hasattr(self.self_attn, "paged_attention") # Self Attention @@ -471,7 +471,7 @@ def LlamaModel_fast_forward( pass # We already handle KV cache position_ids ourselves. - if False:#(past_key_values_length != 0): + if (past_key_values_length != 0): position_ids = torch.arange( past_key_values_length, seq_length + past_key_values_length, dtype = torch.int32, @@ -664,7 +664,7 @@ def LlamaForCausalLM_fast_forward( *args, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: - if causal_mask is None and past_key_values is not None: + if causal_mask is None and past_key_values is None: causal_mask = xformers.attn_bias.LowerTriangularMask() output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions @@ -767,7 +767,7 @@ class FastLlamaModel: @staticmethod def pre_patch(): - # LlamaAttention .forward = LlamaAttention_fast_forward + LlamaAttention .forward = LlamaAttention_fast_forward LlamaSdpaAttention .forward = LlamaAttention_fast_forward LlamaFlashAttention2.forward = LlamaAttention_fast_forward LlamaDecoderLayer .forward = LlamaDecoderLayer_fast_forward