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

Function train

examples/BGRL_G2L.py:115–136  ·  view source on GitHub ↗
(encoder_model, contrast_model, dataloader, optimizer)

Source from the content-addressed store, hash-verified

113
114
115def 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
139def test(encoder_model, dataloader):

Callers 1

mainFunction · 0.70

Calls 1

update_target_encoderMethod · 0.45

Tested by

no test coverage detected