-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathProtea_Predict.py
More file actions
41 lines (36 loc) · 1.8 KB
/
Copy pathProtea_Predict.py
File metadata and controls
41 lines (36 loc) · 1.8 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
import argparse
import os
import Data_Processing as dp
import Pipeline_Functions as pf
def main():
parser = argparse.ArgumentParser(description="Run Protea predictions on PPIGraphs.")
parser.add_argument("--esm2_embeddings_dir", default="./data/esm2_embeddings/", help="Folder containing ESM2 embeddings (.pkl)")
parser.add_argument("--model_address", default="./model/protea_weights.pth", help="Path to Protea weights (.pth)")
parser.add_argument("--raw_data_folder", default="./data/ppigraphs/", help="Folder containing PPIGraph collections (.dat)")
parser.add_argument("--save_folder", default="./data/protea_predictions/", help="Folder to save GNN predictions (.csv)")
parser.add_argument("--batch_size", type=int, default=2, help="Batch size for predictions")
args = parser.parse_args()
# Create save folder if it doesn't exist
if not os.path.exists(args.save_folder):
print(f"Creating missing output directory: {args.save_folder}")
os.makedirs(args.save_folder, exist_ok=True)
print("Loading ESM2 Embeddings...")
esm_dict = {}
pkl_files = [f for f in os.listdir(args.esm2_embeddings_dir) if f.endswith('.pkl')]
if not pkl_files:
print(f"Warning: No ESM2 embedding .pkl files found in {args.esm2_embeddings_dir}")
else:
for file in pkl_files:
address = os.path.join(args.esm2_embeddings_dir, file[:-4])
print(f" - Loading tokens from: {address}")
sub_dict = dp.load_esm_token_dict(address, weighted=True)
esm_dict.update(sub_dict)
pf.run_protea_predictions(
esm_dict=esm_dict,
model_address=args.model_address,
raw_data_folder=args.raw_data_folder,
save_folder=args.save_folder,
batch_size=args.batch_size
)
if __name__ == '__main__':
main()