MCPcopy Create free account
hub / github.com/SLDGroup/EMCAD / get_upsampling_weight

Function get_upsampling_weight

utils/misc.py:28–38  ·  view source on GitHub ↗
(in_channels, out_channels, kernel_size)

Source from the content-addressed store, hash-verified

26
27
28def get_upsampling_weight(in_channels, out_channels, kernel_size):
29 factor = (kernel_size + 1) // 2
30 if kernel_size % 2 == 1:
31 center = factor - 1
32 else:
33 center = factor - 0.5
34 og = np.ogrid[:kernel_size, :kernel_size]
35 filt = (1 - abs(og[0] - center) / factor) * (1 - abs(og[1] - center) / factor)
36 weight = np.zeros((in_channels, out_channels, kernel_size, kernel_size), dtype=np.float64)
37 weight[list(range(in_channels)), list(range(out_channels)), :, :] = filt
38 return torch.from_numpy(weight).float()
39
40
41class CrossEntropyLoss2d(nn.Module):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected