MCPcopy Create free account
hub / github.com/DirectMolecularConfGen/DMCG / evaluate

Function evaluate

train.py:70–96  ·  view source on GitHub ↗
(model, device, loader, args)

Source from the content-addressed store, hash-verified

68
69
70def evaluate(model, device, loader, args):
71 model.eval()
72 mol_labels = []
73 mol_preds = []
74 for batch in tqdm(loader, desc="Iteration", disable=args.disable_tqdm):
75 batch = batch.to(device)
76 with torch.no_grad():
77 pred, _ = model(batch)
78 pred = pred[-1]
79 batch_size = batch.num_graphs
80 n_nodes = batch.n_nodes.tolist()
81 pre_nodes = 0
82 for i in range(batch_size):
83 mol_labels.append(batch.rd_mol[i])
84 mol_preds.append(
85 set_rdmol_positions(batch.rd_mol[i], pred[pre_nodes : pre_nodes + n_nodes[i]])
86 )
87 pre_nodes += n_nodes[i]
88
89 rmsd_list = []
90 for gen_mol, ref_mol in zip(mol_preds, mol_labels):
91 try:
92 rmsd_list.append(get_best_rmsd(gen_mol, ref_mol))
93 except Exception as e:
94 continue
95
96 return np.mean(rmsd_list)
97
98
99def main():

Callers 1

mainFunction · 0.70

Calls 2

set_rdmol_positionsFunction · 0.90
get_best_rmsdFunction · 0.90

Tested by

no test coverage detected