From 9db95ab942da282cd5a13ff145e6b85dfdbb88ae Mon Sep 17 00:00:00 2001 From: George Hotz Date: Mon, 9 Nov 2020 17:56:57 -0800 Subject: [PATCH] fix enet padding --- examples/efficientnet.py | 7 +++++-- 1 file changed, 5 insertions(+), 2 deletions(-) diff --git a/examples/efficientnet.py b/examples/efficientnet.py index a422615f31..cdc85656fa 100644 --- a/examples/efficientnet.py +++ b/examples/efficientnet.py @@ -25,8 +25,11 @@ class MBConvBlock: else: self._expand_conv = None - self.pad = (kernel_size-1)//2 self.strides = strides + if strides == (2,2): + self.pad = [(kernel_size-1)//2-1, (kernel_size-1)//2]*2 + else: + self.pad = [(kernel_size-1)//2]*4 self._depthwise_conv = Tensor.zeros(oup, 1, kernel_size, kernel_size) self._bn1 = BatchNorm2D(oup) @@ -44,7 +47,7 @@ class MBConvBlock: x = inputs if self._expand_conv: x = swish(self._bn0(x.conv2d(self._expand_conv))) - x = x.pad2d(padding=(self.pad, self.pad, self.pad, self.pad)) + x = x.pad2d(padding=self.pad) x = x.conv2d(self._depthwise_conv, stride=self.strides, groups=self._depthwise_conv.shape[0]) x = swish(self._bn1(x))