Download scripts/property_prediction/pdbbind_split.py from OneScience-Group/TargetDiff: direct link, hf CLI and curl.
- Browser
- Download file 2.86 kB
-
https://huggingface.co/OneScience-Group/TargetDiff/resolve/main/scripts/property_prediction/pdbbind_split.py
- Command line
-
hf download hf://OneScience-Group/TargetDiff/scripts/property_prediction/pdbbind_split.py
-
curl -L -o pdbbind_split.py https://huggingface.co/OneScience-Group/TargetDiff/resolve/main/scripts/property_prediction/pdbbind_split.py
2.86 kB
| import os | |
| import pickle | |
| import random | |
| import argparse | |
| import torch | |
| import numpy as np | |
| def coretest_split(index_path, test_path, val_ratio=0.1, val_num=None): | |
| with open(index_path, 'rb') as f: | |
| index = pickle.load(f) | |
| test_ids = [f for f in os.listdir(test_path) if len(f) == 4] | |
| all_ids = [os.path.basename(i[0])[:4] for i in index] | |
| print('Original test set size: ', len(test_ids)) | |
| test_index = [all_ids.index(test_id) for test_id in test_ids if test_id in all_ids] | |
| train_val_index = list(set(range(len(all_ids))) - set(test_index)) | |
| assert len(train_val_index) == len(all_ids) - len(test_index) | |
| random.shuffle(train_val_index) | |
| if val_num is not None: | |
| n_val = val_num | |
| else: | |
| n_val = int(len(train_val_index) * val_ratio) | |
| val_index = train_val_index[:n_val] | |
| train_index = train_val_index[n_val:] | |
| return train_index, val_index, test_index | |
| def time_split(index_path): | |
| valid_ids = np.loadtxt("./data/pdbbind_v2020/timesplit_no_lig_overlap_val", dtype=str) | |
| test_ids = np.loadtxt("./data/pdbbind_v2020/timesplit_test", dtype=str) | |
| with open(index_path, 'rb') as f: | |
| index = pickle.load(f) | |
| all_ids = [os.path.basename(i[0])[:4] for i in index] | |
| val_index = [all_ids.index(val_id) for val_id in valid_ids if val_id in all_ids] | |
| test_index = [all_ids.index(test_id) for test_id in test_ids if test_id in all_ids] | |
| train_index = list(set(range(len(all_ids))) - set(test_index) - set(val_index)) | |
| return train_index, val_index, test_index | |
| if __name__ == '__main__': | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--index_path', type=str, default='./data/pdbbind_v2016/pocket_10/index.pkl') | |
| parser.add_argument('--split_mode', type=str, choices=['coreset', 'time'], default='coreset') | |
| parser.add_argument('--test_path', type=str, default='./data/pdbbind_v2016/coreset') | |
| parser.add_argument('--val_ratio', type=float, default=0.1) | |
| parser.add_argument('--val_num', type=int, default=None) | |
| parser.add_argument('--save_path', type=str, default='./data/pdbbind_v2016/pocket_10/split.pt') | |
| parser.add_argument('--seed', type=int, default=2021) | |
| args = parser.parse_args() | |
| random.seed(args.seed) | |
| if args.split_mode == 'coreset': | |
| train_index, val_index, test_index = coretest_split(args.index_path, args.test_path, | |
| args.val_ratio, args.val_num) | |
| elif args.split_mode == 'time': | |
| train_index, val_index, test_index = time_split(args.index_path) | |
| else: | |
| raise ValueError(args.split_mode) | |
| torch.save({ | |
| 'train': train_index, | |
| 'val': val_index, | |
| 'test': test_index | |
| }, args.save_path) | |
| print('Train %d, Validation %d, Test %d.' % (len(train_index), len(val_index), len(test_index))) | |
| print('Done.') | |