diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index d61dbf842b..53abedf52f 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -124,52 +124,52 @@ def LlamaAttention_fast_forward_inference( # Prefill phase # if not hasattr(self, "paged_attention"): - if do_prefill: - self.paged_attention = torch.empty((KV_CACHE_INCREMENT+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, 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]: - self.paged_attention.resize_((self.paged_attention.shape[0]+KV_CACHE_INCREMENT, 2, bsz, n_kv_heads, head_dim)) - self.paged_attention_K = self.paged_attention[:,0] - self.paged_attention_V = self.paged_attention[:,1] - self.attention.resize_((bsz, n_heads, 1, self.attention.shape[-1]+KV_CACHE_INCREMENT)) - pass - - 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]) + # if do_prefill: + # self.paged_attention = torch.empty((KV_CACHE_INCREMENT+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, 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]: + # self.paged_attention.resize_((self.paged_attention.shape[0]+KV_CACHE_INCREMENT, 2, bsz, n_kv_heads, head_dim)) + # self.paged_attention_K = self.paged_attention[:,0] + # self.paged_attention_V = self.paged_attention[:,1] + # self.attention.resize_((bsz, n_heads, 1, self.attention.shape[-1]+KV_CACHE_INCREMENT)) + # pass + + 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) - # cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len) - # Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids) - cos = self.rotary_emb.cos_cached[seq_len] - sin = self.rotary_emb.sin_cached[seq_len] - h = head_dim // 2 + cos, sin = self.rotary_emb(Vn, seq_len = kv_seq_len) + Qn, Kn = inplace_rope_embedding(Qn, Kn, cos, sin, position_ids) + # 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_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); + # 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 - # Kn = torch.cat([K1, Kn], dim = 2) - # Vn = torch.cat([V1, Vn], dim = 2) - 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) + Kn = torch.cat([K1, Kn], dim = 2) + Vn = torch.cat([V1, Vn], dim = 2) + # 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: @@ -182,13 +182,13 @@ def LlamaAttention_fast_forward_inference( Knn, Vnn = Kn, Vn # Attention - A = torch.matmul(Qn, Knn.transpose(2, 3), out = self.attention[:,:,:,:kv_seq_len]) + 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]) + A = fast_linear_forward(self.o_proj, A)#, out = self.temp_QA[1]) return A, (Kn, Vn) pass @@ -200,13 +200,13 @@ def fast_mlp_inference(self, X): 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]) + gate = fast_linear_forward(self.gate_proj)#, X, out = temp[0]) + up = fast_linear_forward(self. up_proj)#, X, out = temp[1]) gate = torch.nn.functional.silu(gate, inplace = True) gate *= up # X = self.down_proj(gate) - down = fast_linear_forward(self.down_proj, gate, out = up[:,:,:hd]) + down = fast_linear_forward(self.down_proj, gate)#, out = up[:,:,:hd]) return down pass