Commit d076d5c7 authored by Jiaqi Wang's avatar Jiaqi Wang Committed by Kai Chen
Browse files

fix masked_conv cuda runtime error when mask is all zero (#779)

* fix mask conv import error

* fix masked_conv cuda runtime error when mask is all zero
parent 68589d36
...@@ -30,6 +30,8 @@ class MaskedConv2dFunction(Function): ...@@ -30,6 +30,8 @@ class MaskedConv2dFunction(Function):
math.floor((features.size(3) + 2 * pad_w - math.floor((features.size(3) + 2 * pad_w -
(kernel_h - 1) - 1) / stride_w + 1)) (kernel_h - 1) - 1) / stride_w + 1))
mask_inds = torch.nonzero(mask[0] > 0) mask_inds = torch.nonzero(mask[0] > 0)
output = features.new_zeros(batch_size, out_channel, out_h, out_w)
if mask_inds.numel() > 0:
mask_h_idx = mask_inds[:, 0].contiguous() mask_h_idx = mask_inds[:, 0].contiguous()
mask_w_idx = mask_inds[:, 1].contiguous() mask_w_idx = mask_inds[:, 1].contiguous()
data_col = features.new_zeros(in_channel * kernel_h * kernel_w, data_col = features.new_zeros(in_channel * kernel_h * kernel_w,
...@@ -41,7 +43,6 @@ class MaskedConv2dFunction(Function): ...@@ -41,7 +43,6 @@ class MaskedConv2dFunction(Function):
masked_output = torch.addmm(1, bias[:, None], 1, masked_output = torch.addmm(1, bias[:, None], 1,
weight.view(out_channel, -1), data_col) weight.view(out_channel, -1), data_col)
output = features.new_zeros(batch_size, out_channel, out_h, out_w)
masked_conv2d_cuda.masked_col2im_forward(masked_output, mask_h_idx, masked_conv2d_cuda.masked_col2im_forward(masked_output, mask_h_idx,
mask_w_idx, out_h, out_w, mask_w_idx, out_h, out_w,
out_channel, output) out_channel, output)
......
Markdown is supported
0% or .
You are about to add 0 people to the discussion. Proceed with caution.
Finish editing this message first!
Please register or to comment