diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 58459eda56..3b62d4654d 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -308,8 +308,6 @@ def LlamaAttention_fast_forward( padding_mask: Optional[torch.LongTensor] = None, *args, **kwargs, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: - - bsz, q_len, _ = hidden_states.size() # Check for inference if past_key_value is not None: @@ -322,6 +320,8 @@ def LlamaAttention_fast_forward( return A, None, past_key_value pass + bsz, q_len, _ = hidden_states.size() + n_heads = self.num_heads n_groups = self.num_key_value_groups n_kv_heads = self.num_key_value_heads @@ -430,7 +430,6 @@ def LlamaDecoderLayer_fast_forward( (see `past_key_values`). past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states """ - bsz, q_len, hd = hidden_states.size() if past_key_value is not None: # Self Attention residual = hidden_states @@ -658,8 +657,7 @@ def LlamaModel_fast_forward( if output_attentions: all_self_attns += (layer_outputs[1],) pass - - bsz, q_len, hd = hidden_states.size() + if past_key_values is not None: hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states) else: diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index ff082432d1..76f9af11e0 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -46,8 +46,6 @@ def MistralAttention_fast_forward( *args, **kwargs, ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[Tuple[torch.Tensor]]]: - bsz, q_len, _ = hidden_states.size() - # Check for inference if past_key_value is not None: A, past_key_value = LlamaAttention_fast_forward_inference( @@ -59,6 +57,8 @@ def MistralAttention_fast_forward( return A, None, past_key_value pass + bsz, q_len, _ = hidden_states.size() + n_heads = self.num_heads n_groups = self.num_key_value_groups n_kv_heads = self.num_key_value_heads