INNER CODE UNIT · Python

sd

QizhiPei/FABind · FABind/fabind/main_fabind.py:407

        sd = ((com_coord_pred.detach() - com_coord) ** 2).sum(dim=-1)
        rmsd = scatter_mean(sd, index=compound_batch, dim=0).sqrt().detach()

        centroid_pred = scatter_mean(src=com_coord_pred, index=compound_batch, dim=0)
        centroid_true = scatter_mean(src=com_coord, index=compound_batch, dim=0)
        centroid_dis = (centroid_pred - centroid_true).norm(dim=-1)
        
        loss = com_coord_loss + \
            contact_loss + contact_by_pred_loss + contact_distill_loss + \
            pocket_cls_loss + \
            pocket_coord_loss
        
        accelerator.backward(loss)
        if args.clip_grad:
            # clip_grad_norm_(model.parameters(), max_norm=1.0, error_if_nonfinite=True)
            if accelerator.sync_gradients:
                accelerator.clip_grad_norm_(model.parameters(), max_norm=1.0)
        

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…