(encoder_model, contrast_model, dataloader, optimizer)
| 113 | |
| 114 | |
| 115 | def train(encoder_model, contrast_model, dataloader, optimizer): |
| 116 | encoder_model.train() |
| 117 | total_loss = 0 |
| 118 | |
| 119 | for data in dataloader: |
| 120 | data = data.to('cuda') |
| 121 | if data.x is None: |
| 122 | num_nodes = data.batch.size(0) |
| 123 | data.x = torch.ones((num_nodes, 1), dtype=torch.float32).to(data.batch.device) |
| 124 | |
| 125 | optimizer.zero_grad() |
| 126 | _, _, h1_pred, h2_pred, g1_target, g2_target = encoder_model(data.x, data.edge_index, batch=data.batch) |
| 127 | |
| 128 | loss = contrast_model(h1_pred=h1_pred, h2_pred=h2_pred, |
| 129 | g1_target=g1_target.detach(), g2_target=g2_target.detach(), batch=data.batch) |
| 130 | loss.backward() |
| 131 | optimizer.step() |
| 132 | encoder_model.update_target_encoder(0.99) |
| 133 | |
| 134 | total_loss += loss.item() |
| 135 | |
| 136 | return total_loss |
| 137 | |
| 138 | |
| 139 | def test(encoder_model, dataloader): |
no test coverage detected