From 0b661a23b9993e735f74b8be4e973566d8c8df4a Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Fri, 2 Feb 2024 03:55:36 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 25 ++++++++++++++----------- 1 file changed, 14 insertions(+), 11 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 62e06cd7a3..8cfca18755 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -259,17 +259,20 @@ 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) - if not hasattr(self, "paged_attention"): - self.paged_attention = torch.empty((2, bsz, n_kv_heads, 2048, 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 - self.paged_attention_V[:,:,:seq_len,:] = V1 - pass - Kn = self.paged_attention_K[:,:,seq_len,:] = Kn - Vn = self.paged_attention_V[:,:,seq_len,:] = Vn + print(K1.shape, V1.shape) + print(Kn.shape, Vn.shape) + Kn = torch.cat([K1, Kn], dim = 2) + Vn = torch.cat([V1, Vn], dim = 2) + print(Kn.shape, Vn.shape) + # if not hasattr(self, "paged_attention"): + # self.paged_attention = torch.empty((2048, 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 + # self.paged_attention_V[:,:,:seq_len,:] = V1 + # pass + # Kn = self.paged_attention_K[:,:,seq_len,:] = Kn + # Vn = self.paged_attention_V[:,:,seq_len,:] = Vn # Grouped query attention if n_groups != 1: