Views
No views yet

ReactantA.ReactantB>AgentA>ProductA.ProductBO=C([O-])O.[H+]>>O.O=C=O1"""Inference on a SMILES txt. Saved as fastas
2Previously called generate_comparison"""
3
4
5
6def calculate_perplexity(model, src_ids, tgt_ids, tgt_tokenizer):
7 """Conditional perplexity P(tgt | src)."""
8 # make sure padding tokens are ignored
9 labels = tgt_ids.clone()
10 labels[labels == tgt_tokenizer.pad_token_id] = -100
11 with torch.no_grad():
12 # for encoder–decoder models this will set input_ids=src_ids internally
13 outputs = model(input_ids=src_ids, labels=labels)
14 loss = outputs.loss
15 return math.exp(loss.item())
16
17
18if __name__ == '__main__':
19 from transformers import AutoTokenizer, AutoModelForSeq2SeqLM,AutoModelForCausalLM #T5ForConditionalGeneration
20 import argparse
21 import os
22 import torch
23 import json
24 import math
25
26 parser = argparse.ArgumentParser(description='Mol2Pro inference',
27 formatter_class=argparse.ArgumentDefaultsHelpFormatter)
28 parser.add_argument('--input_file', default='../inference/random_smiles2.txt', type=str,
29 help='File with the input molecule SMILES')
30 parser.add_argument('--model_path', default='./output03/checkpoint-60000', type=str, help='Path to model to load')
31 parser.add_argument('--tokenizer_aa',
32 default='/home/woody/b114cb/b114cb10/mol2pro/1.training-different-sizes/1.all-data-16M-tokenizernuria/tokenizer_aa', type=str,
33 help='Path to amino acid tokenizer')
34 parser.add_argument('--tokenizer_mol',
35 default='/home/woody/b114cb/b114cb10/mol2pro/1.training-different-sizes/1.all-data-16M-tokenizernuria/nuria_tokenizer_smiles', type=str,
36 help='Path to SMILES tokenizer')
37 parser.add_argument('--top_k',
38 default=100,type=int,
39 help='K for top-k sampling')
40 parser.add_argument('--output_folder', default='fastas', type=str, help='Folder for saving results')
41 parser.add_argument('--top_p', default =1.0, type=float)
42 parser.add_argument('--repetition_penalty', default=1.0, type=float)
43
44 args = parser.parse_args()
45
46 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
47
48 if 'gatgpt' in args.model_path.lower():
49 GNN = True
50 print('Graph data mode')
51 else:
52 GNN = False
53 print('SMILES/SELFIES data mode')
54
55 # Load protein tokenizer
56 if 'ape' in args.tokenizer_aa:
57 from ape_tokenizer import APETokenizer
58
59 tokenizer_aa = APETokenizer.from_pretrained(args.tokenizer_aa)
60 else:
61 tokenizer_aa = AutoTokenizer.from_pretrained(args.tokenizer_aa)
62
63 # Load molecule tokenizer
64 if GNN:
65 tokenizer_mol = None
66 else:
67 if 'ape' in args.tokenizer_mol:
68 from ape_tokenizer import APETokenizer
69
70 tokenizer_mol = APETokenizer.from_pretrained(args.tokenizer_mol)
71 else:
72 tokenizer_mol = AutoTokenizer.from_pretrained(args.tokenizer_mol)
73
74 # Load model
75 dec_only = False
76 if GNN:
77 from transformers import GPT2Config, Trainer
78 from models import GATGPT2Config, GATGPT2
79 from torch_geometric.data import Batch, Data
80
81 config = GATGPT2Config.from_pretrained(args.model_path)
82
83 # Load model weights
84 model = GATGPT2.from_pretrained(args.model_path, config=config)
85
86 model.eval()
87 model.to("cuda" if torch.cuda.is_available() else "cpu")
88 else:
89 try:
90 print('Attempt Seq2Seq model load... ')
91 model = AutoModelForSeq2SeqLM.from_pretrained(args.model_path).cuda()
92 except:
93 print('Attempt CausalLM model load... ')
94 model = AutoModelForCausalLM.from_pretrained(args.model_path).cuda()
95 model.config.eos_token_id = tokenizer_mol.eos_token_id
96 model.config.pad_token_id = tokenizer_mol.pad_token_id
97 print(
98 f"Set `eos_token_id` to {tokenizer_mol.eos_token_id} and `pad_token_id` to {tokenizer_mol.pad_token_id}.")
99 dec_only = True
100 print('Model Loaded')
101
102
103 smiles_list = []
104 with open(args.input_file, 'r') as input_file:
105 for line in input_file:
106 smiles_list.append(line.strip())
107
108 molecule_json = {}
109 for index,smiles in enumerate(smiles_list):
110 sequences=[]
111 if GNN:
112 from build_tokenized_dataset import convert_smiles_to_graph
113
114 node_feats, edge_index, edge_feats = convert_smiles_to_graph(smiles)
115 node_feats_tensor = torch.tensor(node_feats, dtype=torch.float, device=device)
116 edge_index_tensor = torch.tensor(edge_index, dtype=torch.long, device=device).T.contiguous()
117 edge_feats_tensor = torch.tensor(edge_feats, dtype=torch.float, device=device)
118
119 # Input to decoder is only bos
120 start_token = tokenizer_aa.bos_token_id or tokenizer_aa.convert_tokens_to_ids("▁") # fallback to the space which is always appended by our tokenizer
121 text_input_ids = torch.tensor([[start_token]], dtype=torch.long, device=device)
122
123 input_ids = {
124 "graph_node_feats": node_feats_tensor, # shape (N, 3)
125 "graph_edge_index": edge_index_tensor, # shape (2, E)
126 "graph_edge_feats": edge_feats_tensor, # shape (E, 2)
127 "batch": torch.full((len(node_feats),), 0, dtype=torch.long, device=device), # shape (N,)
128 "input_ids": text_input_ids
129 }
130
131 elif 'ape' in args.tokenizer_mol:
132 input_ids = tokenizer_mol(smiles, return_tensors="pt")["input_ids"].to(device='cuda')
133 else:
134 input_ids = tokenizer_mol(smiles, return_tensors="pt").input_ids.to(device='cuda')
135 if not GNN:
136 print(f'Generating for {smiles} (input ids: {input_ids})')
137 else:
138 print(f'Generating for {smiles}')
139
140 # top_k = Choose at random from the first K tokens (weigthed by softmax score)
141 # num_return_sequences = The number of independently computed returned sequences for each element in the batch.
142 if dec_only:
143 attention_mask = torch.ones_like(input_ids).cuda()
144 outputs = model.generate(input_ids, top_k=args.top_k, top_p=args.top_p, attention_mask = attention_mask, repetition_penalty=args.repetition_penalty, max_length=1024, do_sample=True, num_return_sequences=25)
145 else:
146 outputs = model.generate(input_ids, top_k=args.top_k, top_p=args.top_p, repetition_penalty=args.repetition_penalty, max_length=1024, do_sample=True, num_return_sequences=25)
147
148 sequences = [tokenizer_aa.decode(output, skip_special_tokens=True) for output in outputs]
149
150 ppls = []
151 for out_ids in outputs:
152 tgt = out_ids.unsqueeze(0)
153 p = calculate_perplexity(model,src_ids=input_ids,tgt_ids=tgt,tgt_tokenizer=tokenizer_aa)
154 ppls.append(p)
155
156 if not os.path.exists(args.output_folder):
157 os.makedirs(args.output_folder)
158
159 filename = f'{args.output_folder}/output_topk{args.top_k}_file-{index}.fasta'
160 with open(filename, 'w') as fn:
161 for idx, (seq, ppl) in enumerate(zip(sequences, ppls)):
162 fn.write(f">{idx}|ppl={ppl:.2f}\n")
163 fn.write(seq + "\n")
164
165 # Store molecule name
166 molecule_json[filename] = smiles
167
168 # Save metadata
169 metadata_path = os.path.join(args.output_folder, 'molecule_input_metadata.json')
170 try:
171 with open(metadata_path, 'w') as json_file:
172 json.dump(molecule_json, json_file, indent=4)
173 print(f"Metadata successfully written to {metadata_path}")
174 except Exception as e:
175 print(f"An error occurred while writing to JSON: {e}")1INFERENCE_FOLDER=output_folder # change to the name of the output folder you want
2MODEL=checkpoint-90000 # path to the model
3INFERENCE_TXT=reaction.txt # text file containing the reactions (in SMILE format) wanting to generate for.
4
5REPETITION_PENALTY=1.0
6TOP_P=1.0
7TOP_K=100
8
9source .environment/bin/activate # load an environment containing the required dependencies (transformers, torch, datasets)
10python inference.py --input_file "$INFERENCE_TXT" --model_path "$MODEL" --output_folder "$INFERENCE_FOLDER" --tokenizer_mol tokenizer_ABPE_rexzyme_offset --tokenizer_aa tokenizer_aa