INNER CODE UNIT · Python
post_optim_mol
QizhiPei/FABind · FABind/fabind/fabind_inference.py:285
def post_optim_mol(args, accelerator, data, com_coord_pred, com_coord_pred_per_sample_list, com_coord_per_sample_list, compound_batch, LAS_tmp, rigid=False):
post_optim_device='cpu'
for i in range(compound_batch.max().item()+1):
i_mask = (compound_batch == i)
com_coord_pred_i = com_coord_pred[i_mask]
com_coord_i = data[i]['compound'].rdkit_coords
com_coord_pred_center_i = com_coord_pred_i.mean(dim=0).reshape(1, 3)
if rigid:
predict_coord, loss, rmsd = post_optimize_compound_coords(
reference_compound_coords=com_coord_i.to(post_optim_device),
predict_compound_coords=com_coord_pred_i.to(post_optim_device),
LAS_edge_index=None,
mode=args.post_optim_mode,
total_epoch=args.post_optim_epoch,
)
predict_coord.to(accelerator.device)