MCPcopy Create free account
hub / github.com/PyGCL/PyGCL / Encoder

Class Encoder

examples/BGRL_L2L.py:60–100  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

58
59
60class Encoder(torch.nn.Module):
61 def __init__(self, encoder, augmentor, hidden_dim, dropout=0.2, predictor_norm='batch'):
62 super(Encoder, self).__init__()
63 self.online_encoder = encoder
64 self.target_encoder = None
65 self.augmentor = augmentor
66 self.predictor = torch.nn.Sequential(
67 torch.nn.Linear(hidden_dim, hidden_dim),
68 Normalize(hidden_dim, norm=predictor_norm),
69 torch.nn.PReLU(),
70 torch.nn.Dropout(dropout))
71
72 def get_target_encoder(self):
73 if self.target_encoder is None:
74 self.target_encoder = copy.deepcopy(self.online_encoder)
75
76 for p in self.target_encoder.parameters():
77 p.requires_grad = False
78 return self.target_encoder
79
80 def update_target_encoder(self, momentum: float):
81 for p, new_p in zip(self.get_target_encoder().parameters(), self.online_encoder.parameters()):
82 next_p = momentum * p.data + (1 - momentum) * new_p.data
83 p.data = next_p
84
85 def forward(self, x, edge_index, edge_weight=None):
86 aug1, aug2 = self.augmentor
87 x1, edge_index1, edge_weight1 = aug1(x, edge_index, edge_weight)
88 x2, edge_index2, edge_weight2 = aug2(x, edge_index, edge_weight)
89
90 h1, h1_online = self.online_encoder(x1, edge_index1, edge_weight1)
91 h2, h2_online = self.online_encoder(x2, edge_index2, edge_weight2)
92
93 h1_pred = self.predictor(h1_online)
94 h2_pred = self.predictor(h2_online)
95
96 with torch.no_grad():
97 _, h1_target = self.get_target_encoder()(x1, edge_index1, edge_weight1)
98 _, h2_target = self.get_target_encoder()(x2, edge_index2, edge_weight2)
99
100 return h1, h2, h1_pred, h2_pred, h1_target, h2_target
101
102
103def train(encoder_model, contrast_model, data, optimizer):

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected