Update llama.py

This commit is contained in:
Daniel Han-Chen 2024-02-01 22:48:54 +11:00
commit 334c5ed1f0

View file

@ -235,9 +235,9 @@ def LlamaAttention_fast_forward_inference(
temp_QA = torch.empty((2, bsz, 1, hd), dtype = dtype, device = "cuda")
temp_KV = torch.empty((2, bsz, 1, n_kv_heads*head_dim), dtype = dtype, device = "cuda")
Qn = fast_linear_forward(self.q_proj, Xn)#, out = temp_QA[0])
Kn = fast_linear_forward(self.k_proj, Xn)#, out = temp_KV[0])
Vn = fast_linear_forward(self.v_proj, Xn)#, out = temp_KV[1])
Qn = fast_linear_forward(self.q_proj, Xn, out = temp_QA[0])
Kn = fast_linear_forward(self.k_proj, Xn, out = temp_KV[0])
Vn = fast_linear_forward(self.v_proj, Xn, out = 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)
@ -279,7 +279,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 = temp_QA[1])
A = fast_linear_forward(self.o_proj, A, out = temp_QA[1])
return A, (Kn, Vn)
pass
@ -291,13 +291,13 @@ def fast_mlp_inference(self, X):
mlp_size = self.config.intermediate_size
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