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