From e094914b0ea2a13ed1cfbcd5f5dd7b940e383dce Mon Sep 17 00:00:00 2001 From: Daniel Han-Chen Date: Thu, 8 Feb 2024 02:30:49 +1100 Subject: [PATCH] Update llama.py --- unsloth/models/llama.py | 18 +++++++++++++----- 1 file changed, 13 insertions(+), 5 deletions(-) diff --git a/unsloth/models/llama.py b/unsloth/models/llama.py index e38137cab7..a7385be3cc 100644 --- a/unsloth/models/llama.py +++ b/unsloth/models/llama.py @@ -206,21 +206,29 @@ def LlamaAttention_fast_forward_inference( return A, (Kn, Vn) pass - +@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 fast_mlp_inference(self, X): # gate = self.gate_proj(X) # up = self.up_proj(X) bsz, _, hd = X.shape mlp_size = self.config.intermediate_size - temp = torch.empty((2, bsz, 1, mlp_size), dtype = X.dtype, device = "cuda") + #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