MCPcopy Create free account
hub / github.com/SwinTransformer/Transformer-SSL / MoBY

Class MoBY

models/moby.py:32–167  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

30
31
32class MoBY(nn.Module):
33 def __init__(self,
34 cfg,
35 encoder,
36 encoder_k,
37 contrast_momentum=0.99,
38 contrast_temperature=0.2,
39 contrast_num_negative=4096,
40 proj_num_layers=2,
41 pred_num_layers=2,
42 **kwargs):
43 super().__init__()
44
45 self.cfg = cfg
46
47 self.encoder = encoder
48 self.encoder_k = encoder_k
49
50 self.contrast_momentum = contrast_momentum
51 self.contrast_temperature = contrast_temperature
52 self.contrast_num_negative = contrast_num_negative
53
54 self.proj_num_layers = proj_num_layers
55 self.pred_num_layers = pred_num_layers
56
57 self.projector = MoBYMLP(in_dim=self.encoder.num_features, num_layers=proj_num_layers)
58 self.projector_k = MoBYMLP(in_dim=self.encoder.num_features, num_layers=proj_num_layers)
59 self.predictor = MoBYMLP(num_layers=pred_num_layers)
60
61 for param_q, param_k in zip(self.encoder.parameters(), self.encoder_k.parameters()):
62 param_k.data.copy_(param_q.data) # initialize
63 param_k.requires_grad = False # not update by gradient
64
65 for param_q, param_k in zip(self.projector.parameters(), self.projector_k.parameters()):
66 param_k.data.copy_(param_q.data)
67 param_k.requires_grad = False
68
69 if self.cfg.MODEL.SWIN.NORM_BEFORE_MLP == 'bn':
70 nn.SyncBatchNorm.convert_sync_batchnorm(self.encoder)
71 nn.SyncBatchNorm.convert_sync_batchnorm(self.encoder_k)
72
73 nn.SyncBatchNorm.convert_sync_batchnorm(self.projector)
74 nn.SyncBatchNorm.convert_sync_batchnorm(self.projector_k)
75 nn.SyncBatchNorm.convert_sync_batchnorm(self.predictor)
76
77 self.K = int(self.cfg.DATA.TRAINING_IMAGES * 1. / dist.get_world_size() / self.cfg.DATA.BATCH_SIZE) * self.cfg.TRAIN.EPOCHS
78 self.k = int(self.cfg.DATA.TRAINING_IMAGES * 1. / dist.get_world_size() / self.cfg.DATA.BATCH_SIZE) * self.cfg.TRAIN.START_EPOCH
79
80 # create the queue
81 self.register_buffer("queue1", torch.randn(256, self.contrast_num_negative))
82 self.register_buffer("queue2", torch.randn(256, self.contrast_num_negative))
83 self.queue1 = F.normalize(self.queue1, dim=0)
84 self.queue2 = F.normalize(self.queue2, dim=0)
85
86 self.register_buffer("queue_ptr", torch.zeros(1, dtype=torch.long))
87
88 @torch.no_grad()
89 def _momentum_update_key_encoder(self):

Callers 1

build_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected