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']

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…