rxn-sandbox / scripts /predict_retrosynthesis_tree.py
helderlopes's picture
Initial commit
25f9bfc
Raw
History Blame Contribute Delete
2.13 kB
"""
Update this script and run it inside the worker container to run your retrosynthesis tree predictions via script
"""
print("Setting up retrosynthesis tree prediction...")
# 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)
# 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)