(self, x, edge_index, edge_weight=None)
| 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 | |
| 103 | def train(encoder_model, contrast_model, data, optimizer): |
nothing calls this directly
no test coverage detected