MCPcopy Create free account
hub / github.com/QWTforGithub/T2LDM / _make_fuse_layers

Method _make_fuse_layers

timm/models/hrnet.py:442–476  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

440 return nn.ModuleList(branches)
441
442 def _make_fuse_layers(self):
443 if self.num_branches == 1:
444 return nn.Identity()
445
446 num_branches = self.num_branches
447 num_inchannels = self.num_inchannels
448 fuse_layers = []
449 for i in range(num_branches if self.multi_scale_output else 1):
450 fuse_layer = []
451 for j in range(num_branches):
452 if j > i:
453 fuse_layer.append(nn.Sequential(
454 nn.Conv2d(num_inchannels[j], num_inchannels[i], 1, 1, 0, bias=False),
455 nn.BatchNorm2d(num_inchannels[i], momentum=_BN_MOMENTUM),
456 nn.Upsample(scale_factor=2 ** (j - i), mode='nearest')))
457 elif j == i:
458 fuse_layer.append(nn.Identity())
459 else:
460 conv3x3s = []
461 for k in range(i - j):
462 if k == i - j - 1:
463 num_outchannels_conv3x3 = num_inchannels[i]
464 conv3x3s.append(nn.Sequential(
465 nn.Conv2d(num_inchannels[j], num_outchannels_conv3x3, 3, 2, 1, bias=False),
466 nn.BatchNorm2d(num_outchannels_conv3x3, momentum=_BN_MOMENTUM)))
467 else:
468 num_outchannels_conv3x3 = num_inchannels[j]
469 conv3x3s.append(nn.Sequential(
470 nn.Conv2d(num_inchannels[j], num_outchannels_conv3x3, 3, 2, 1, bias=False),
471 nn.BatchNorm2d(num_outchannels_conv3x3, momentum=_BN_MOMENTUM),
472 nn.ReLU(False)))
473 fuse_layer.append(nn.Sequential(*conv3x3s))
474 fuse_layers.append(nn.ModuleList(fuse_layer))
475
476 return nn.ModuleList(fuse_layers)
477
478 def get_num_inchannels(self):
479 return self.num_inchannels

Callers 1

__init__Method · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected