From d6ab9c92d7bd7a71eae2420091da7e43c8f5cb80 Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Thu, 8 Feb 2024 02:49:26 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 17 +++++++++++++---- 1 file changed, 13 insertions(+), 4 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index a7385be3cc..698193b2d1 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -74,6 +74,15 @@ pass from math import sqrt as math_sqrt KV_CACHE_INCREMENT = 128 # KV Cache update size +@torch.compile(options = { + "epilogue_fusion" : True, + "max_autotune" : True, + "fallback_random" : False, + "shape_padding" : True, + "triton.cudagraphs" : False, + "trace.enabled" : True, + "trace.graph_diagram" : True, +}, dynamic = True,) def LlamaAttention_fast_forward_inference( self, hidden_states: torch.Tensor, @@ -142,9 +151,9 @@ def LlamaAttention_fast_forward_inference( 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 = 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) @@ -202,7 +211,7 @@ def LlamaAttention_fast_forward_inference( 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