INNER CODE UNIT · Python
batch_len
QizhiPei/FABind · FABind/fabind/main_fabind.py:450
batch_len = protein_out_mask_whole.sum(dim=1).detach()
protein_len_list.append(batch_len)
pocket_coord_pred_list.append(pred_pocket_center.detach())
pocket_coord_list.append(data.coords_center)
# use hard to calculate acc and skip samples
for i, j in enumerate(batch_len):
count += 1
pocket_cls_list.append(pocket_cls.detach()[i][:j])
pocket_cls_pred_list.append(pocket_cls_pred.detach()[i][:j].sigmoid())
pocket_cls_pred_round_list.append(pocket_cls_pred.detach()[i][:j].sigmoid().round().int())
pred_index_bool = (pocket_cls_pred.detach()[i][:j].sigmoid().round().int() == 1)
if pred_index_bool.sum() == 0: # all the prediction is False, skip
skip_count += 1
if batch_id % args.log_interval == 0:
stats_dict = {}
stats_dict['step'] = batch_id
stats_dict['lr'] = optimizer.param_groups[0]['lr']