diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 6ed52a7acb..30348b69c2 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -705,26 +705,24 @@ 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 \ - hasattr(self.model.layers[0].self_attn, "paged_attention"): + if past_key_values is not None and hasattr(self.model.layers[0].self_attn, "paged_attention"): outputs = fast_forward_inference( self.model, input_ids, 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 +986,7 @@ class FastLlamaModel: padding_side = "right", token = token, ) - + model, tokenizer = patch_tokenizer(model, tokenizer) model = model_patcher.post_patch(model) diff --git a/unsloth/save.py b/unsloth/save.py index 5971d76e6e..42d326e128 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -90,14 +90,14 @@ def _merge_lora(layer, name): W = fast_dequantize(W, quant_state) else: dtype = W.dtype - W = W.to(torch.float32).t() - # W = W.t() + # W = W.to(torch.float32).t() + W = W.t() if A is not None: # sAB = (A.t().to(torch.float32) @ (s * B.t().to(torch.float32))) # W += sAB - W.addmm_(A.t().to(torch.float32), B.t().to(torch.float32), alpha = s) - # W.addmm_(A.t().to(W.dtype), B.t().to(W.dtype), alpha = s) + # W.addmm_(A.t().to(torch.float32), B.t().to(torch.float32), alpha = s) + W.addmm_(A.t().to(W.dtype), B.t().to(W.dtype), alpha = s) # if not torch.isfinite(W).all(): maximum_element = torch.max(W.min().abs(), W.max()) if not torch.isfinite(maximum_element).item():