INNER CODE UNIT · Python
centroid_dis
QizhiPei/FABind · FABind/fabind/main_fabind.py:412
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)
optimizer.step()
scheduler.step()
batch_loss += len(y_pred)*contact_loss.item()
batch_by_pred_loss += len(y_pred_by_coord)*contact_by_pred_loss.item()