diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 44de7bf391..f0e615b5a1 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -250,7 +250,7 @@ def LlamaAttention_fast_forward_inference( sin = self.rotary_emb.sin_cached[seq_len] h = head_dim // 2 - RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = Xn.dtype, device = "cuda") + RH_Q = torch.empty((bsz, n_heads, 1, head_dim), dtype = dtype, device = "cuda") 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); @@ -259,8 +259,17 @@ def LlamaAttention_fast_forward_inference( Kn *= cos; Kn.addcmul_(RH_K, sin); # New KV cache - Kn = torch.cat([K1, Kn], dim = 2) - Vn = torch.cat([V1, Vn], dim = 2) + # Kn = torch.cat([K1, Kn], dim = 2) + # Vn = torch.cat([V1, Vn], dim = 2) + if not hasattr(self, "paged_attention_K"): + paged_attention = torch.empty((2, bsz, n_kv_heads, 2048, head_dim), dtype = dtype, device = "cuda") + self.paged_attention_K = paged_attention[0] + self.paged_attention_V = paged_attention[1] + self.paged_attention_K[:,:,:kv_seq_len,:] = K1 + self.paged_attention_V[:,:,:kv_seq_len,:] = V1 + pass + Kn = self.paged_attention_K[:,:,seq_len,:] = K1 + Vn = self.paged_attention_V[:,:,seq_len,:] = V1 # Grouped query attention if n_groups != 1: