From 9c2bed35b966e622390da3363cc27e637e2d30b1 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 2 Feb 2024 23:45:02 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 122 +++++++--------------------------------- 1 file changed, 21 insertions(+), 101 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index e7cffbb7b6..80537f70fd 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -72,11 +72,12 @@ pass from math import sqrt as math_sqrt -def LlamaAttention_fast_forward_inference_prefill( +def LlamaAttention_fast_forward_inference( self, hidden_states: torch.Tensor, past_key_value: Optional[Tuple[torch.Tensor]], position_ids, + do_prefill = False, ): """ https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L406 @@ -121,16 +122,18 @@ def LlamaAttention_fast_forward_inference_prefill( # Prefill phase # if not hasattr(self, "paged_attention"): - self.paged_attention = torch.empty((self.config.max_position_embeddings+1, 2, bsz, n_kv_heads, head_dim), dtype = dtype, device = "cuda") - self.paged_attention_K = self.paged_attention[:,0] - 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, self.config.max_position_embeddings), dtype = dtype, device = "cuda") - self.scalar = 1.0 / math_sqrt(self.head_dim) + if do_prefill: + self.paged_attention = torch.empty((self.config.max_position_embeddings+1, 2, bsz, n_kv_heads, head_dim), dtype = dtype, device = "cuda") + self.paged_attention_K = self.paged_attention[:,0] + 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, self.config.max_position_embeddings), dtype = dtype, device = "cuda") + self.scalar = 1.0 / math_sqrt(self.head_dim) + pass # pass Qn = fast_linear_forward(self.q_proj, Xn, out = self.temp_QA[0]) @@ -184,71 +187,6 @@ def LlamaAttention_fast_forward_inference_prefill( pass -def LlamaAttention_fast_forward_inference( - self, - hidden_states: torch.Tensor, - past_key_value: Optional[Tuple[torch.Tensor]], - position_ids, -): - Xn = hidden_states - bsz, _, hd = hidden_states.size() - K1, V1 = past_key_value - - n_heads = self.num_heads - n_groups = self.num_key_value_groups - n_kv_heads = self.num_key_value_heads - head_dim = self.head_dim - seq_len = K1.shape[-2] - kv_seq_len = seq_len + 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) - - # RoPE - cos = self.rotary_emb.cos_cached[seq_len] - sin = self.rotary_emb.sin_cached[seq_len] - h = head_dim // 2 - - 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); - - RH_K = RH_Q[:,:n_kv_heads,:,:] # torch.empty((n_kv_heads, 1, head_dim), dtype = dtype, device = "cuda") - RH_K[:,:,:,:h] = Kn[:,:,:,h:]; RH_K[:,:,:,h:] = Kn[:,:,:,:h]; torch.neg(RH_K[:,:,:,:h], out = RH_K[:,:,:,:h]); - Kn *= cos; Kn.addcmul_(RH_K, sin); - - # New KV cache - self.paged_attention_K[seq_len] = Kn.permute(2, 0, 1, 3) - self.paged_attention_V[seq_len] = Vn.permute(2, 0, 1, 3) - Kn = self.paged_attention_K[:kv_seq_len].permute(1, 2, 0, 3) - Vn = self.paged_attention_V[:kv_seq_len].permute(1, 2, 0, 3) - - # Grouped query attention - if n_groups != 1: - _, _, cached_len, _ = Kn.shape - Knn = Kn[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, cached_len, head_dim) - Vnn = Vn[:, :, None, :, :].expand(bsz, n_kv_heads, n_groups, cached_len, head_dim) - Knn = Knn.reshape(bsz, n_heads, cached_len, head_dim) - Vnn = Vnn.reshape(bsz, n_heads, cached_len, head_dim) - else: - Knn, Vnn = Kn, Vn - - # Attention - A = torch.matmul(Qn, Knn.transpose(2, 3), out = self.attention[:,:,:,:kv_seq_len]) - A *= self.scalar - A[:] = torch.nn.functional.softmax(A, dim = -1, dtype = torch.float32)#.to(A.dtype) - 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 = self.temp_QA[1]) - return A, (Kn, Vn) -pass - - def fast_mlp_inference(self, X): # gate = self.gate_proj(X) # up = self.up_proj(X) @@ -415,7 +353,9 @@ def LlamaDecoderLayer_fast_forward( (see `past_key_values`). past_key_value (`Tuple(torch.FloatTensor)`, *optional*): cached past key and value projection states """ - if past_key_value is not None and hasattr(self.self_attn, "paged_attention"): + if past_key_value is not None: + do_prefill = not hasattr(self.self_attn, "paged_attention") + # Self Attention residual = hidden_states hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states) @@ -424,23 +364,7 @@ def LlamaDecoderLayer_fast_forward( hidden_states, past_key_value, position_ids, - ) - hidden_states += residual - - # Fully Connected - residual = hidden_states - hidden_states = fast_rms_layernorm_inference(self.post_attention_layernorm, hidden_states) - hidden_states = fast_mlp_inference(self.mlp, hidden_states) - hidden_states += residual - elif past_key_value is not None: - # Self Attention - residual = hidden_states - hidden_states = fast_rms_layernorm_inference(self.input_layernorm, hidden_states) - hidden_states, present_key_value = LlamaAttention_fast_forward_inference_prefill( - self.self_attn, - hidden_states, - past_key_value, - position_ids, + do_prefill = do_prefill, ) hidden_states += residual @@ -655,12 +579,8 @@ def LlamaModel_fast_forward( if output_attentions: all_self_attns += (layer_outputs[1],) pass - - if past_key_values is not None: - hidden_states = fast_rms_layernorm_inference(self.norm, hidden_states) - else: - hidden_states = fast_rms_layernorm(self.norm, hidden_states) - pass + + hidden_states = fast_rms_layernorm(self.norm, hidden_states) # add hidden states from the last decoder layer if output_hidden_states: @@ -679,7 +599,7 @@ pass # https://github.com/huggingface/transformers/blob/main/src/transformers/models/llama/modeling_llama.py#L825 -@torch.compile +@torch.inference_mode def LlamaModel_fast_forward_inference( self, input_ids,