diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 237ec6cc1c..5c16a050cc 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -78,10 +78,8 @@ def LlamaAttention_fast_forward_inference( self, hidden_states: torch.Tensor, past_key_value: Optional[Tuple[torch.Tensor]], + position_ids, do_prefill = False, - temp_QA = None, - temp_KV = None, - RH_Q = None, ): """ https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L406 @@ -132,6 +130,9 @@ def LlamaAttention_fast_forward_inference( self.paged_attention_V = self.paged_attention[:,1] self.paged_attention_K[:seq_len] = K1.permute(2, 0, 1, 3) self.paged_attention_V[:seq_len] = V1.permute(2, 0, 1, 3) + self.temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda") + self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda") + self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda") self.attention = torch.empty((bsz, n_heads, 1, KV_CACHE_INCREMENT), dtype = dtype, device = "cuda") self.scalar = 1.0 / math_sqrt(self.head_dim) elif kv_seq_len >= self.paged_attention.shape[0]: @@ -141,15 +142,9 @@ def LlamaAttention_fast_forward_inference( self.attention.resize_((bsz, n_heads, 1, self.attention.shape[-1]+KV_CACHE_INCREMENT)) pass - if temp_QA is None: - temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda") - temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda") - RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda") - pass - - Qn = fast_linear_forward(self.q_proj, Xn, out = temp_QA[0]) - Kn = fast_linear_forward(self.k_proj, Xn, out = temp_KV[0]) - Vn = fast_linear_forward(self.v_proj, Xn, out = temp_KV[1]) + Qn = fast_linear_forward(self.q_proj, Xn, out = self.temp_QA[0]) + Kn = fast_linear_forward(self.k_proj, Xn, out = self.temp_KV[0]) + Vn = fast_linear_forward(self.v_proj, Xn, out = self.temp_KV[1]) Qn = Qn.view(bsz, 1, n_heads, head_dim).transpose(1, 2) Kn = Kn.view(bsz, 1, n_kv_heads, head_dim).transpose(1, 2) Vn = Vn.view(bsz, 1, n_kv_heads, head_dim).transpose(1, 2) @@ -160,7 +155,7 @@ def LlamaAttention_fast_forward_inference( sin = self.rotary_emb.sin_cached[seq_len] h = head_dim // 2 - # RH_Q = self.RH_Q + RH_Q = self.RH_Q RH_Q[:,:,:,:h] = Qn[:,:,:,h:]; RH_Q[:,:,:,h:] = Qn[:,:,:,:h]; torch.neg(RH_Q[:,:,:,:h], out = RH_Q[:,:,:,:h]); Qn *= cos; Qn.addcmul_(RH_Q, sin); @@ -193,19 +188,17 @@ def LlamaAttention_fast_forward_inference( A = torch.matmul(A, Vnn, out = Qn) A = A.transpose(1, 2) A = A.reshape(bsz, 1, self.hidden_size) - A = fast_linear_forward(self.o_proj, A, out = temp_QA[1]) + A = fast_linear_forward(self.o_proj, A, out = self.temp_QA[1]) return A, (Kn, Vn) pass -def fast_mlp_inference(self, X, temp = None): +def fast_mlp_inference(self, X): # gate = self.gate_proj(X) # up = self.up_proj(X) bsz, _, hd = X.shape - if temp is None: - mlp_size = self.config.intermediate_size - temp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda") - pass + mlp_size = self.config.intermediate_size + temp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda") gate = fast_linear_forward(self.gate_proj, X, out = temp[0]) up = fast_linear_forward(self. up_proj, X, out = temp[1]) @@ -218,26 +211,13 @@ def fast_mlp_inference(self, X, temp = None): pass -def fast_rms_layernorm_inference(self, X, temp1 = None, temp2 = None): +def fast_rms_layernorm_inference(self, X): old_dtype = X.dtype - - if temp1 is None: - bsz, _, hd = X.shape - temp1 = torch.empty((2, bsz, 1, hd), dtype = torch.float32, device = "cuda") - temp2 = torch.empty((bsz, 1, hd), dtype = old_dtype, device = "cuda") - pass - - XX = temp1[0] - XX2 = temp1[1] - - # XX = X.to(torch.float32) - XX[:] = X - # variance = XX.square().mean(-1, keepdim = True) - variance = torch.square(XX, out = XX2).mean(-1, keepdim = True) + XX = X.to(torch.float32) + variance = XX.square().mean(-1, keepdim = True) variance += self.variance_epsilon XX *= variance.rsqrt_() - # X = XX.to(old_dtype) # Must preserve due to residual - temp2[:] = XX; X = temp2 + X = XX.to(old_dtype) # Must preserve due to residual X *= self.weight return X pass @@ -262,6 +242,9 @@ def LlamaAttention_fast_forward( del self.paged_attention_K del self.paged_attention_V del self.paged_attention + del self.temp_QA + del self.temp_KV + del self.RH_Q del self.attention pass @@ -385,6 +368,7 @@ def LlamaDecoderLayer_fast_forward( self.self_attn, hidden_states, past_key_value, + position_ids, do_prefill = do_prefill, ) hidden_states += residual @@ -443,16 +427,6 @@ def LlamaModel_fast_forward( return_dict: Optional[bool] = None, *args, **kwargs, ) -> Union[Tuple, BaseModelOutputWithPast]: - - # Clear inference - if hasattr(self, "temp_QA"): - del self.temp_QA - del self.temp_KV - del self.RH_Q - del self.mlp_temp - del self.layernorm1 - del self.layernorm2 - pass output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions assert(output_attentions is False) @@ -638,78 +612,31 @@ def LlamaModel_fast_forward_inference( ): # Fix out of bounds tokenization input_ids = input_ids[:,:self.max_seq_length] + hidden_states = self.embed_tokens(input_ids) - bsz, q_len, hd = hidden_states.shape - dtype = hidden_states.dtype - - first_attention = self.layers[0].self_attn - n_heads = first_attention.num_heads - n_groups = first_attention.num_key_value_groups - n_kv_heads = first_attention.num_key_value_heads - head_dim = first_attention.head_dim - mlp_size = self.config.intermediate_size - - # Temporary matrices - if not hasattr(self, "temp_QA"): - self.temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda") - self.temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda") - self.RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda") - self.mlp_temp = torch.empty((2, bsz, 1, mlp_size), dtype = dtype, device = "cuda") - self.layernorm1 = torch.empty((2, bsz, 1, hd), dtype = torch.float32, device = "cuda") - self.layernorm2 = torch.empty((bsz, 1, hd), dtype = dtype, device = "cuda") - pass - temp_QA = self.temp_QA - temp_KV = self.temp_KV - RH_Q = self.RH_Q - mlp_temp = self.mlp_temp - layernorm1 = self.layernorm1 - layernorm2 = self.layernorm2 - # next_decoder_cache = [] for idx, decoder_layer in enumerate(self.layers): # Self Attention residual = hidden_states - hidden_states = fast_rms_layernorm_inference( - decoder_layer.input_layernorm, - hidden_states, - layernorm1, - layernorm2, - ) + hidden_states = fast_rms_layernorm_inference(decoder_layer.input_layernorm, hidden_states) hidden_states, present_key_value = LlamaAttention_fast_forward_inference( decoder_layer.self_attn, hidden_states, past_key_values[idx], - do_prefill = False, - temp_QA = temp_QA, - temp_KV = temp_KV, - RH_Q = RH_Q, + None, ) hidden_states += residual # Fully Connected residual = hidden_states - hidden_states = fast_rms_layernorm_inference( - decoder_layer.post_attention_layernorm, - hidden_states, - layernorm1, - layernorm2, - ) - hidden_states = fast_mlp_inference( - decoder_layer.mlp, - hidden_states, - mlp_temp, - ) + hidden_states = fast_rms_layernorm_inference(decoder_layer.post_attention_layernorm, hidden_states) + hidden_states = fast_mlp_inference(decoder_layer.mlp, hidden_states) hidden_states += residual next_decoder_cache.append(present_key_value) pass - hidden_states = fast_rms_layernorm_inference( - self.norm, - hidden_states, - layernorm1, - layernorm2, - ) + hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states) return BaseModelOutputWithPast( last_hidden_state = hidden_states, diff --git a/unsloth/models/mistral.py b/unsloth/models/mistral.py index bf6b3e1f78..14e77dfb8e 100644 --- a/unsloth/models/mistral.py +++ b/unsloth/models/mistral.py @@ -51,6 +51,9 @@ def MistralAttention_fast_forward( del self.paged_attention_K del self.paged_attention_V del self.paged_attention + del self.temp_QA + del self.temp_KV + del self.RH_Q del self.attention pass