INNER CODE UNIT · Python
get_data
QizhiPei/FABind · FABind/fabind/data.py:98
def get_data(args, logger, addNoise=None, use_whole_protein=False, compound_coords_init_mode='pocket_center_rdkit', pre="/PDBbind_data/pdbbind2020"):
if args.data == "0":
logger.log_message(f"Loading dataset")
logger.log_message(f"compound feature based on torchdrug")
logger.log_message(f"protein feature based on esm2")
add_noise_to_com = float(addNoise) if addNoise else None
new_dataset = FABindDataSet(f"{pre}/dataset", add_noise_to_com=add_noise_to_com, use_whole_protein=use_whole_protein, compound_coords_init_mode=compound_coords_init_mode, pocket_radius=args.pocket_radius, noise_for_predicted_pocket=args.noise_for_predicted_pocket,
test_random_rotation=args.test_random_rotation, pocket_idx_no_noise=args.pocket_idx_no_noise, use_esm2_feat=args.use_esm2_feat, seed=args.seed, pre=pre, args=args)
# load compound features extracted using torchdrug.
# c_length: number of atoms in the compound
# This filter may cause some samples to be filtered out. So the actual number of samples is less than that in the original papers.
train_tmp = new_dataset.data.query("c_length < 100 and native_num_contact > 5 and group =='train' and use_compound_com").reset_index(drop=True)
valid_test_tmp = new_dataset.data.query("(group == 'valid' or group == 'test') and use_compound_com").reset_index(drop=True)
new_dataset.data = pd.concat([train_tmp, valid_test_tmp], axis=0).reset_index(drop=True)
d = new_dataset.data
only_native_train_index = d.query("group =='train'").index.values
train = new_dataset[only_native_train_index]