-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmain.py
More file actions
184 lines (152 loc) · 8.74 KB
/
Copy pathmain.py
File metadata and controls
184 lines (152 loc) · 8.74 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
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
import torch
import torch.nn.functional as F
import logging
import argparse
import pickle
from tqdm import tqdm
from logger.logger import create_logger
from loader.loader import create_loader
from models.master import create_model
from train.train import train
from torch_geometric.graphgym.utils.comp_budget import params_count
from torch_geometric import seed_everything
from torch_geometric.graphgym.config import cfg, set_cfg
from torch_geometric.graphgym.logger import set_printing
def inference(model, loader):
"""
Run inference using the trained model and data loader, compute metrics, and save the results.
Args:
model (torch.nn.Module): The trained model to be evaluated.
loader (torch.utils.data.DataLoader): DataLoader for the dataset to perform inference on.
This function sets the model to evaluation mode and disables gradient calculations.
It iterates over the data loader, collects predictions and ground truths, and computes metrics such as IoU,
Mean Absolute Error (MAE), and similarity index for each batch. The metrics are logged, and all inference outputs
are saved to a pickle file specified by `cfg.inference_output`.
"""
from train.metrics import compute_3D_IoU, get_similarity_index
model.eval()
with torch.no_grad():
inference_output = {"pred": [], "true": [], "temp": [], "cell": [], "refcode": [], "pos": [], "atoms": [], "iou": [], "mae": [], "similarity_index": []}
for iter, batch in tqdm(enumerate(loader), total=len(loader), ncols=50):
batch.to("cuda:0")
inference_output["cell"].append(batch.cell.detach().to("cpu"))
inference_output["atoms"].append(batch.x[batch.non_H_mask].detach().to("cpu"))
inference_output["pos"].append(batch.pos[batch.non_H_mask].detach().to("cpu"))
inference_output["refcode"].append(batch.refcode[0])
inference_output["temp"].append(batch.temperature_og.detach().to("cpu")[0])
_pred, _true = model(batch)
inference_output["pred"].append(_pred.detach().to("cpu"))
inference_output["true"].append(_true.detach().to("cpu"))
inference_output["iou"].append(compute_3D_IoU(_pred, _true).detach().to("cpu"))
inference_output["mae"].append(F.l1_loss(_pred,_true, reduce="none").detach().to("cpu"))
inference_output["similarity_index"].append(get_similarity_index(_pred, _true).detach().to("cpu"))
iou = torch.cat(inference_output["iou"], dim=0)
mae = torch.cat(inference_output["mae"], dim=0)
similarity_index = torch.cat(inference_output["similarity_index"], dim=0)
logging.info(f"Mean IoU: {iou.mean().item()} +/- {iou.std().item()}")
logging.info(f"Mean MAE: {mae.mean().item()} +/- {mae.std().item()}")
logging.info(f"Mean Similarity Index: {similarity_index.mean().item()} +/- {similarity_index.std().item()}")
pickle.dump(inference_output, open(cfg.inference_output, "wb"))
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument('--seed', type=int, default=0, help='Seed for the experiment')
parser.add_argument('--name', type=str, default="CartNet", help="name of the Wandb experiment" )
parser.add_argument("--batch", type=int, default=4, help="Batch size")
parser.add_argument("--batch_accumulation", type=int, default=16, help="Batch Accumulation")
parser.add_argument("--dataset", type=str, default="jarvis", help="Dataset name. Available: jarvis, megnet, matbench")
parser.add_argument("--dataset_path", type=str, default="./dataset/jarvis/")
parser.add_argument("--inference", action="store_true", help="Inference")
parser.add_argument("--checkpoint_path", type=str, default=None, help="Path of the checkpoints of the model")
parser.add_argument("--inference_output", type=str, default="./inference.pkl", help="Path to the inference output")
parser.add_argument("--figshare_target", type=str, default="formation_energy_peratom", help="Figshare dataset target")
parser.add_argument("--wandb_project", type=str, default="PRISM_crystal", help="Wandb project name")
parser.add_argument("--wandb_entity", type=str, default="PRISM_PROJECT", help="Name of the wandb entity")
parser.add_argument("--loss", type=str, default="MAE", help="Loss function")
parser.add_argument("--epochs", type=int, default=500, help="Number of epochs")
parser.add_argument("--lr", type=float, default=1e-3, help="Learning rate")
parser.add_argument("--warmup", type=float, default=0.01, help="Warmup")
parser.add_argument('--model', type=str, default="PRISM", help="Model Name")
parser.add_argument("--optimizer", type=str, default="adam_schedulefree", help="Optimizer")
parser.add_argument("--w_decay", type=float, default=0.0, help="Weight decay")
parser.add_argument("--beta1", type=float, default=0.9, help="Beta 1")
parser.add_argument("--epsilon", type=float, default=1e-8, help="Epsilon")
parser.add_argument("--radius", type=float, default=5.0, help="Radius for the Radius Graph Neighbourhood")
parser.add_argument("--rbf_feat", type=float, default=5.0, help="RBF feature")
parser.add_argument("--radius_cell", type=float, default=15.0, help="Radius for the Radius Graph Neighbourhood")
parser.add_argument("--radius_feat", type=float, default=1.0, help="Radius feature")
parser.add_argument("--num_layers", type=int, default=4, help="Number of layers")
parser.add_argument("--dim_in", type=int, default=256, help="Input dimension")
parser.add_argument("--dim_rbf", type=int, default=64, help="Number of RBF")
parser.add_argument("--dropout", type=float, default=0.0, help="Dropout")
parser.add_argument('--augment', action='store_true', help='augment')
parser.add_argument("--invariant", action="store_true", help="Rotation Invariant model")
parser.add_argument("--disable_envelope", action="store_false", help="Disable envelope")
parser.add_argument('--disable_atom_types', action='store_false', help='Atom types')
parser.add_argument("--workers", type=int, default=5, help="Number of workers")
set_cfg(cfg)
args = parser.parse_args()
cfg.seed = args.seed
cfg.name = args.name
cfg.run_dir = "results/"+cfg.name+"/"+str(cfg.seed)
cfg.inference_output = args.inference_output
cfg.dataset.task_type = "regression"
cfg.batch = args.batch
cfg.batch_accumulation = args.batch_accumulation
cfg.dataset.name = args.dataset
cfg.dataset_path = args.dataset_path
cfg.figshare_target = args.figshare_target
cfg.wandb_project = args.wandb_project
cfg.wandb_entity = args.wandb_entity
cfg.loss = args.loss
cfg.optim.max_epoch = args.epochs
cfg.lr = args.lr
cfg.warmup = args.warmup
cfg.model = args.model
cfg.radius = args.radius
cfg.rbf_feat = args.rbf_feat
cfg.radius_cell = args.radius_cell
cfg.radius_feat = args.radius_feat
cfg.optimizer = args.optimizer
cfg.w_decay = args.w_decay
cfg.beta1 = args.beta1
cfg.epsilon = args.epsilon
cfg.num_layers = args.num_layers
cfg.dim_in = args.dim_in
cfg.dim_rbf = args.dim_rbf
cfg.augment = args.augment
cfg.invariant = args.invariant
cfg.envelope = args.disable_envelope
cfg.use_atom_types = args.disable_atom_types
cfg.workers = args.workers
cfg.dropout = args.dropout
set_printing()
#Seed
seed_everything(cfg.seed)
logging.info(f"Experiment will be saved at: {cfg.run_dir}")
loaders = create_loader()
model = create_model()
logging.info(model)
cfg.params_count = params_count(model)
logging.info(f"Number of parameters: {cfg.params_count}")
if cfg.optimizer == "adam_schedulefree":
import schedulefree
warmup_steps = int(cfg.optim.max_epoch * len(loaders[0]) * cfg.warmup)
optimizer = schedulefree.AdamWScheduleFree(model.parameters(),
lr=cfg.lr,
weight_decay=cfg.w_decay,
betas=(args.beta1, 0.999),
eps=args.epsilon,
warmup_steps=warmup_steps)
elif cfg.optimizer == "adam":
optimizer = torch.optim.Adam(model.parameters(), lr=cfg.lr)
else:
raise Exception("Optimizer not implemented")
loggers = create_logger()
if args.inference:
assert args.checkpoint_path is not None, "Weights path not provided"
ckpt = torch.load(args.checkpoint_path)
model.load_state_dict(ckpt["model_state"])
cfg.inference_output = args.inference_output
inference(model, loaders[-1])
else:
train(model, loaders, optimizer, loggers)