From c0d95162556757d4e0b688574bb36aeebe07cf18 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Sun, 10 Mar 2024 13:10:03 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 27 +++++++++++++-------------- 1 file changed, 13 insertions(+), 14 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 6ed52a7acb..ae52312deb 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -705,19 +705,7 @@ def CausalLM_fast_forward(fast_forward_inference): *args, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: - 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 - output_hidden_states = ( - output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states - ) - return_dict = return_dict if return_dict is not None else self.config.use_return_dict - - # decoder outputs consists of (dec_features, layer_state, dec_hidden, dec_attn) - self.model._has_no_labels = labels is None - - if past_key_values is not None and \ + if False: #past_key_values is not None and \ hasattr(self.model.layers[0].self_attn, "paged_attention"): outputs = fast_forward_inference( self.model, @@ -725,6 +713,17 @@ def CausalLM_fast_forward(fast_forward_inference): past_key_values, ) else: + causal_mask = xformers.attn_bias.LowerTriangularMask() + + output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + # 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, @@ -988,7 +987,7 @@ class FastLlamaModel: padding_side = "right", token = token, ) - + model, tokenizer = patch_tokenizer(model, tokenizer) model = model_patcher.post_patch(model)