From dfe9f9889c9f08923d50c5b4fc8b6122066307d2 Mon Sep 17 00:00:00 2001 From: Ming Liu Date: Wed, 26 Feb 2020 14:03:21 +0800 Subject: [PATCH 1/2] Optimize MPNCOV.py for less memory usage and faster inference --- TestCode/code/model/MPNCOV/python/MPNCOV.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/TestCode/code/model/MPNCOV/python/MPNCOV.py b/TestCode/code/model/MPNCOV/python/MPNCOV.py index 9683f8d..cc815cd 100644 --- a/TestCode/code/model/MPNCOV/python/MPNCOV.py +++ b/TestCode/code/model/MPNCOV/python/MPNCOV.py @@ -21,9 +21,10 @@ def forward(ctx, input): w = x.data.shape[3] M = h*w x = x.reshape(batchSize,dim,M) - I_hat = (-1./M/M)*torch.ones(M,M,device = x.device) + (1./M)*torch.eye(M,M,device = x.device) - I_hat = I_hat.view(1,M,M).repeat(batchSize,1,1).type(x.dtype) - y = x.bmm(I_hat).bmm(x.transpose(1,2)) + I_hat = torch.empty(M, M, device=x.device).fill_(-1./M/M) + I_hat_diag = I_hat.diagonal() + I_hat_diag += (1./M) + y = (x @ I_hat).bmm(x.transpose(1,2)) ctx.save_for_backward(input,I_hat) return y @staticmethod From 4378317736607a19966cf7c0e8f79edf1a7b935a Mon Sep 17 00:00:00 2001 From: Ming Liu Date: Thu, 19 Mar 2020 23:36:43 +0800 Subject: [PATCH 2/2] Fix backward function Since I_hat is not a three-dimensioni tensor now, the bmm in the backward function should be replaced by @ as well --- TestCode/code/model/MPNCOV/python/MPNCOV.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/TestCode/code/model/MPNCOV/python/MPNCOV.py b/TestCode/code/model/MPNCOV/python/MPNCOV.py index cc815cd..eccd2f3 100644 --- a/TestCode/code/model/MPNCOV/python/MPNCOV.py +++ b/TestCode/code/model/MPNCOV/python/MPNCOV.py @@ -24,7 +24,7 @@ def forward(ctx, input): I_hat = torch.empty(M, M, device=x.device).fill_(-1./M/M) I_hat_diag = I_hat.diagonal() I_hat_diag += (1./M) - y = (x @ I_hat).bmm(x.transpose(1,2)) + y = x @ I_hat @ x.transpose(1,2) ctx.save_for_backward(input,I_hat) return y @staticmethod @@ -38,7 +38,7 @@ def backward(ctx, grad_output): M = h*w x = x.reshape(batchSize,dim,M) grad_input = grad_output + grad_output.transpose(1,2) - grad_input = grad_input.bmm(x).bmm(I_hat) + grad_input = grad_input @ x @ I_hat grad_input = grad_input.reshape(batchSize,dim,h,w) return grad_input