From 1ba28d8e26ee3d46e4b2f47e59eb65a216f604f9 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Tue, 23 Jan 2024 01:33:51 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 23 ++++++++++++++--------- 1 file changed, 14 insertions(+), 9 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index 763c669286..9d1df56680 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -139,7 +139,7 @@ def LlamaAttention_fast_forward_inference( 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 = torch.empty((n_kv_heads, head_dim), dtype = dtype, device = "cuda") + RH_K = RH_Q[:n_kv_heads, :] # torch.empty((n_kv_heads, head_dim), dtype = dtype, device = "cuda") RH_K[:, :, :h] = Qn[:, :, h:]; RH_K[:, :, h:] = Kn[:, :, :h]; torch.neg(RH_K[:, :, :h], out = RH_K[:, :, :h]); Kn *= cos; Kn.addcmul_(RH_K, sin); @@ -149,11 +149,15 @@ def LlamaAttention_fast_forward_inference( # Grouped query attention if n_groups != 1: - _, _, cached_len, _ = Kn.shape - Knn = Kn[:, :, None, :, :].expand(1, n_kv_heads, n_groups, cached_len, head_dim) - Vnn = Vn[:, :, None, :, :].expand(1, n_kv_heads, n_groups, cached_len, head_dim) - Knn = Knn.reshape(1, n_heads, cached_len, head_dim) - Vnn = Vnn.reshape(1, n_heads, cached_len, head_dim) + # _, _, cached_len, _ = Kn.shape + # Knn = Kn[:, :, None, :, :].expand(1, n_kv_heads, n_groups, cached_len, head_dim) + # Vnn = Vn[:, :, None, :, :].expand(1, n_kv_heads, n_groups, cached_len, head_dim) + # Knn = Knn.reshape(1, n_heads, cached_len, head_dim) + # Vnn = Vnn.reshape(1, n_heads, cached_len, head_dim) + Knn = Kn[:, :, None, :, :].expand(n_kv_heads, n_groups, kv_seq_len, head_dim) + Vnn = Vn[:, :, None, :, :].expand(n_kv_heads, n_groups, kv_seq_len, head_dim) + Knn = Knn.reshape(n_heads, kv_seq_len, head_dim) + Vnn = Vnn.reshape(n_heads, kv_seq_len, head_dim) else: Knn, Vnn = Kn, Vn @@ -175,7 +179,8 @@ pass torch_silu = torch.nn.functional.silu def fast_mlp_inference(self, X): - X = X.view(self.hidden_size) + hidden_size = self.hidden_size + X = X.view(hidden_size) # gate = self.gate_proj(X) # up = self.up_proj(X) @@ -185,8 +190,8 @@ def fast_mlp_inference(self, X): gate *= up # X = self.down_proj(gate) - down = fast_linear_forward(self.down_proj, gate) - X = down.view(1, 1, self.hidden_size) + down = fast_linear_forward(self.down_proj, gate, out = up[:hidden_size]) + X = down.view(1, 1, hidden_size) return X pass