backward speed slower than pytorch conv2d

#13 · closed · 1 comments

View on GitHub ↗

ybzhou

as mentioned in #11 code in following ``` import torch import time from torch import nn from torch.autograd import Variable from pyinn.modules import Conv2dDepthwise in_feat = 32 out_feat = 64 batch_size = 256 stride = (2,1) kernel_size = (7,3) x = Variable(torch.randn(batch_size, in_feat, 40, 400)).cuda() t = Variable(torch.randn(batch_size)).cuda() print(kernel_size, stride) layers = [nn.BatchNorm2d(in_feat), nn.Conv2d(in_feat, in_feat, kernel_size=kernel_size, stride=stride, padding=0, groups=in_feat, bias=False), nn.BatchNorm2d(in_feat), nn.Conv2d(in_feat, out_feat, kernel_size=(1,1), stride=(1,1), bias=False ), nn.BatchNorm2d(out_feat), nn.ReLU()] torch_conv = nn.Sequential(*layers) torch_conv = torch_conv.cuda() layers = [nn.BatchNorm2d(in_feat), Conv2dDepthwise(in_feat, kernel_size=kernel_size, stride=stride, padding=0, bias=False), nn.BatchNorm2d(in_feat), nn.Conv2d(in_feat, out_feat, kernel_size=(1,1), stride=(1,1), bias=False ), nn.BatchNorm2d(out_feat), nn.ReLU()] pyinn_conv = nn.Sequential(*layers) pyinn_conv = pyinn_conv.cuda() print('forward1') t_start = time.time() for i in range(10): y1 = torch_conv(x) torch.cuda.synchronize() pytorch_forward_t = time.time()-t_start y1 = y1.sum(1).sum(1).sum(1) loss1 = nn.MSELoss()(y1, t) print('backward1') t_start = time.time() loss1.backward() pytorch_backward_t = time.time() - t_start print('forward2') t_start = time.time() for i in range(10): y2 = pyinn_conv(x) torch.cuda.synchronize() pyinn_forward_t = time.time()-t_start y2 = y2.sum(1).sum(1).sum(1) loss2 = nn.MSELoss()(y2, t) print('backward2') t_start = time.time() loss2.backward() pyinn_backward_t = time.time() - t_start print('batch size: {}, kernel size: {}, stride size: {}, ' 'pytorch forward time: {:.4f}s ' 'pyinn forward time: {:.4f}s ' 'pytorch backward time: {:.4f}s ' 'pyinn backward time: {:.4f}s'.format( batch_size, kernel_size, stride, pytorch_forward_t, pyinn_forward_t, pytorch_backward_t, pyinn_backward_t )) print(pytorch_forward_t/pyinn_forward_t, pytorch_backward_t/pyinn_backward_t) ``` ``` batch size: 256, kernel size: (7, 3), stride size: (2, 1), pytorch forward time: 2.4811s pyinn forward time: 1.5070s pytorch backward time: 0.0037s pyinn backward time: 0.1597s 1.646389619890211 0.022934876056741653 ```

Comments

szagoruyko

reopened #11