inference
This commit is contained in:
parent
f41a437540
commit
085a8e944a
2 changed files with 1 additions and 2 deletions
|
|
@ -180,7 +180,6 @@ pass
|
|||
def fast_linear_forward(proj, X, temp_lora = None, out = None):
|
||||
W, W_quant, lora_A, lora_B, lora_S = get_lora_parameters(proj)
|
||||
out = fast_gemv(X, W, W_quant, out = out)
|
||||
print(X.shape, W.quant_state.shape, out.shape)
|
||||
if lora_A is not None:
|
||||
dtype = X.dtype
|
||||
temp_lora = torch.matmul(X, lora_A.to(dtype).t(), out = temp_lora)
|
||||
|
|
|
|||
|
|
@ -172,7 +172,7 @@ def fast_mlp_inference(self, X):
|
|||
gate *= up
|
||||
|
||||
# X = self.down_proj(gate)
|
||||
down = fast_linear_forward(self.down_proj, X)
|
||||
down = fast_linear_forward(self.down_proj, gate)
|
||||
X = down.view(1, 1, self.hidden_size)
|
||||
|
||||
return X
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue