(model, device, loader, args)
| 68 | |
| 69 | |
| 70 | def 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 | |
| 99 | def main(): |
no test coverage detected