rxn-sandbox / scripts /run_notebook_examples.py
helderlopes's picture
Initial commit
25f9bfc
Raw
History Blame Contribute Delete
6.05 kB
print("Setup underway...")
# All necessary imports
from celery import Celery
from utils import wait_for_result
# Initialize Celery
celery_app = Celery()
print("Broker:", celery_app.conf.broker_url)
print("Backend:", celery_app.conf.result_backend)
print("")
print("============================")
print("1a. Product prediction - batch")
# Set up a list of reactants to make predictions
reactants_list = ["CCI.O=Cc1ccc([N+](=O)[O-])c(O)c1", "CCOc1cc(O)c(C=O)cc1OCC.OCCCBr", "C=CCc1cc(OCc2ccccc2)ccc1O.CCBr"]
# Setup task kwargs
kwargs = {
"topn": 1, # Number of results per reactant
"num_beams": 3, # Number of beams used for prediction. Must be >= topn
"device": None, # Device used for predicting, either "cuda" or "cpu", None defaults to cuda if available
"ckpt_forward": "Pistachio2025Q2-Forward", # Default forward model
"vocab": "Pistachio2025Q2", # Vocab for default forward model
# "ckpt_forward_path": "models/forward/Pistachio2025Q2-Forward.ckpt", # Can be used instead of ckpt_forward
# "vocab_path": "vocab/Pistachio2025Q2.txt", # Can be used instead of vocab
}
# Send the product_prediction task with the reaction list and kwargs
task = celery_app.send_task(
"tasks.product_prediction",
[reactants_list],
kwargs=kwargs,
queue="product_prediction",
)
print("Task sent. Assigned task_id: {}".format(task.id))
# Use the task id to get the result. Increase timeout if needed.
wait_for_result(celery_app, task.id, timeout=180)
print("")
print("============================")
print("1b. Product prediction - top 3")
# Set up a list of reactants to make predictions
reactants_list = ["CCI.O=Cc1ccc([N+](=O)[O-])c(O)c1"]
# Setup task kwargs
kwargs = {
"topn": 3, # Number of results per reactant
"num_beams": 5, # Number of beams used for prediction. Must be >= topn
"device": None, # Device used for predicting, either "cuda" or "cpu", None defaults to cuda if available
"ckpt_forward": "Pistachio2025Q2-Forward", # Default forward model
"vocab": "Pistachio2025Q2", # Vocab for default forward model
# "ckpt_forward_path": "models/forward/Pistachio2025Q2-Forward.ckpt", # Can be used instead of ckpt_forward
# "vocab_path": "vocab/Pistachio2025Q2.txt", # Can be used instead of vocab
}
# Send the product_prediction task with the reaction list and kwargs
task = celery_app.send_task(
"tasks.product_prediction",
[reactants_list],
kwargs=kwargs,
queue="product_prediction",
)
print("Task sent. Assigned task_id: {}".format(task.id))
# Use the task id to get the result. Increase timeout if needed.
wait_for_result(celery_app, task.id, timeout=180)
print("")
print("============================")
print("2. Retrosynthesis prediction")
# Choose product for retrosynthesis prediction
product = "C=CC(=C)C[Si](C)(C)C"
# Setup task kwargs
kwargs = {
"topn": 15, # Number of results per reactant
"num_beams": 15, # Number of beams used for prediction. Must be >= topn
"fap": 0.6, # Forward likelihood acceptance probability (not length averaged)
"fld": 0.2, # Forward likelihood delta required between the top2 forward prediction results
"device": None, # Device used for predicting, either "cuda" or "cpu", None defaults to cuda if available
"ckpt_forward": "Pistachio2025Q2-Forward", # Default forward model
"ckpt_retro": "Pistachio2025Q2-Retro", # Default retrosynthesis model
"vocab": "Pistachio2025Q2", # Vocab for default forward and retrosynthesis models
# "ckpt_forward_path": "models/forward/Pistachio2025Q2-Forward.ckpt", # Can be used instead of ckpt_forward
# "ckpt_retro_path": "models/retrosynthesis/Pistachio2025Q2-Retro.ckpt", # Can be used instead of ckpt_retro
# "vocab_path": "vocab/Pistachio2025Q2.txt", # Can be used instead of vocab
}
# Send the retro_prediction task with the product and kwargs
task = celery_app.send_task(
"tasks.retro_prediction",
[product],
kwargs=kwargs,
queue="retro_prediction",
)
print("Task sent. Assigned task_id: {}".format(task.id))
# Use the task id to get the result. Increase timeout if needed.
wait_for_result(celery_app, task.id, timeout=300)
print("")
print("============================")
print("3. Retro tree prediction")
# Choose product for retrosynthesis tree prediction
product = "C1C(C[Si](C)(C)C)=CCC2C(=O)OC(=O)C12"
#product = "Cc1cc2scnc2cc1N"
#product = "CC(C)(C(=O)O)C1C=CC=C(C2CC2)C1=O"
#product = "Nc1ccc2scnc2c1Br"
# Setup task kwargs
kwargs = {
"topn": 15, # Number of results per reactant
"num_beams": 15, # Number of beams used for prediction. Must be >= topn
"fap": 0.6, # Forward likelihood acceptance probability (not length averaged)
"fld": 0.2, # Forward likelihood delta required between the top2 forward prediction results
"max_depth": 4, # Max depth of the retrosynthesis tree
"beam_width": 6, # Max amount of nodes being expanded in each step
"device": None, # Device used for predicting, either "cuda" or "cpu", None defaults to cuda if available
"ckpt_forward": "Pistachio2025Q2-Forward", # Default forward model
"ckpt_retro": "Pistachio2025Q2-Retro", # Default retrosynthesis model
"vocab": "Pistachio2025Q2", # Vocab for default forward and retrosynthesis models
# "ckpt_forward_path": "models/forward/Pistachio2025Q2-Forward.ckpt", # Can be used instead of ckpt_forward
# "ckpt_retro_path": "models/retrosynthesis/Pistachio2025Q2-Retro.ckpt", # Can be used instead of ckpt_retro
# "vocab_path": "vocab/Pistachio2025Q2.txt", # Can be used instead of vocab
}
# Send the retro_prediction_tree task with the product and kwargs
task = celery_app.send_task(
"tasks.retro_prediction_tree",
[product],
kwargs=kwargs,
queue="retro_prediction",
)
print("Task sent. Assigned task_id: {}".format(task.id))
# Use the task id to get the result. Increase timeout if needed.
wait_for_result(celery_app, task.id, timeout=3000)