-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
43 lines (36 loc) · 1.5 KB
/
Copy pathtrain.py
File metadata and controls
43 lines (36 loc) · 1.5 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
import torch
import torch.nn as nn
import torch.nn.functional as F
def train_one_epoch(model, loader, optimizer, device):
model.train()
total_loss = 0.0
for orig_batch in loader:
optimizer.zero_grad()
y_pred = model(
x_f=orig_batch['x'].to(device),
edge_index=orig_batch['edge_index'].to(device),
edge_attr=orig_batch['edge_attr'].to(device),
batch=orig_batch['batch'].to(device)
)
y_true = orig_batch['y'].to(device)
loss = F.mse_loss( y_pred, y_true, reduction='mean')
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
total_loss += loss.item()
return total_loss / len(loader)
def evaluate(model, loader, device):
model.eval()
preds_out, labels_out = [], []
with torch.no_grad():
for orig_batch in loader:
y_pred = model(
x_f=orig_batch['x'].to(device),
edge_index=orig_batch['edge_index'].to(device),
edge_attr=orig_batch['edge_attr'].to(device),
batch=orig_batch['batch'].to(device)
)
y_true = orig_batch['y'].to(device)
preds_out += y_pred.detach().float().cpu().view(-1).tolist()
labels_out += y_true.detach().float().cpu().view(-1).tolist()
return labels_out, preds_out