diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index f24ecafbe3..4d2b03a00f 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -309,7 +309,6 @@ def LlamaAttention_fast_forward( V = V.transpose(1, 2) A = flash_attn_func(Q, K, V, causal = True) else: - print("0", end = "") # Grouped query attention if n_groups != 1: K = K[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim) @@ -515,6 +514,7 @@ def LlamaModel_fast_forward( # if 0 in attention_mask: # padding_mask = attention_mask # else: + print(attention_mask) padding_mask = None attention_mask = _prepare_4d_causal_attention_mask_for_sdpa( @@ -1208,11 +1208,6 @@ class FastLlamaModel: @staticmethod def for_inference(model): - if not hasattr(model, "_original_forward"): - model._original_forward = model.forward - pass - model.forward = torch.inference_mode(model._original_forward) - internal_model = model internal_model.gradient_checkpointing = False internal_model.training = False @@ -1227,10 +1222,6 @@ class FastLlamaModel: @staticmethod def for_training(model, use_gradient_checkpointing = True): - if hasattr(model, "_original_forward"): - model.forward = model._original_forward - pass - internal_model = model internal_model.gradient_checkpointing = use_gradient_checkpointing internal_model.training = True diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index 4cfeb4ad98..bc00e7a982 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -128,7 +128,7 @@ def MistralAttention_fast_forward( A = xformers_attention(Q, K, V, attn_bias = causal_mask) A = A.view(bsz, q_len, n_heads, head_dim) - elif (HAS_FLASH_ATTENTION and attention_mask is None): + elif HAS_FLASH_ATTENTION and attention_mask is None: Q = Q.transpose(1, 2) K = K.transpose(1, 2) V = V.transpose(1, 2) @@ -137,7 +137,6 @@ def MistralAttention_fast_forward( window = (-1, -1) if (kv_seq_len <= sw) else (sw, sw) A = flash_attn_func(Q, K, V, causal = True, window_size = window) else: - print("0", end = "") # Grouped query attention # if n_groups != 1: K = K[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, kv_seq_len, head_dim) @@ -178,7 +177,7 @@ def MistralForCausalLM_fast_forward( *args, **kwargs, ) -> Union[Tuple, CausalLMOutputWithPast]: - if causal_mask is None: + if causal_mask is None and past_key_values is None: bsz, q_len = input_ids.shape sliding_window = getattr(self.config, "sliding_window", None) if sliding_window is None or sliding_window == "null" or sliding_window <= 0: