diff --git a/maxsmi/full_workflow.py b/maxsmi/full_workflow.py index b50801b..dfcb788 100644 --- a/maxsmi/full_workflow.py +++ b/maxsmi/full_workflow.py @@ -50,7 +50,9 @@ NB_EPOCHS, ) -if __name__ == "__main__": + +def main(): + warnings.filterwarnings("ignore") parser = argparse.ArgumentParser() @@ -432,3 +434,7 @@ ) results_metrics = results_metrics.to_pickle(f"{folder}/results_metrics.pkl") logging.info("Script completed. \n \n") + + +if __name__ == "__main__": + main() diff --git a/maxsmi/full_workflow_earlystopping.py b/maxsmi/full_workflow_earlystopping.py index 33d52ca..17eec40 100644 --- a/maxsmi/full_workflow_earlystopping.py +++ b/maxsmi/full_workflow_earlystopping.py @@ -50,7 +50,9 @@ NB_EPOCHS, ) -if __name__ == "__main__": + +def main(): + warnings.filterwarnings("ignore") parser = argparse.ArgumentParser() @@ -451,3 +453,7 @@ ) results_metrics = results_metrics.to_pickle(f"{folder}/results_metrics.pkl") logging.info("Script completed. \n \n") + + +if __name__ == "__main__": + main() diff --git a/maxsmi/prediction_unlabeled_data.py b/maxsmi/prediction_unlabeled_data.py index 6af40d0..266bc96 100644 --- a/maxsmi/prediction_unlabeled_data.py +++ b/maxsmi/prediction_unlabeled_data.py @@ -2,6 +2,7 @@ From smiles to predictions """ +from pathlib import Path import argparse import logging import logging.handlers @@ -36,7 +37,11 @@ from maxsmi.pytorch_evaluation import out_of_sample_prediction from maxsmi.utils_optimal_model import retrieve_optimal_model -if __name__ == "__main__": +PATH_MAXSMI = Path(__file__).parent + + +def main(): + warnings.filterwarnings("ignore") parser = argparse.ArgumentParser() @@ -170,7 +175,7 @@ (ml_model_name, ml_model) = model_type(ml_model, device, smi_dict, max_length_smi) logging.info(f"Summary of ml model used for the prediction: {ml_model} ") - file_path = f"maxsmi/prediction_models/{args.task}" + file_path = PATH_MAXSMI / f"prediction_models/{args.task}" ml_model.load_state_dict( torch.load(f"{file_path}/model_dict.pth", map_location=device) ) @@ -246,3 +251,7 @@ logging.info("Script completed. \n \n") print(f"Script completed. Output can be found at {folder}/") + + +if __name__ == "__main__": + main() diff --git a/setup.py b/setup.py index 970c737..49e5a0e 100644 --- a/setup.py +++ b/setup.py @@ -38,6 +38,13 @@ # Customize MANIFEST.in if the general case does not suit your needs # Comment out this line to prevent the files from being packaged with your software include_package_data=True, + entry_points={ + "console_scripts": [ + "maxsmi = maxsmi.full_workflow:main", + "maxsmi-earlystopping = maxsmi.full_workflow_earlystopping:main", + "maxsmi-pred = maxsmi.prediction_unlabeled_data:main", + ] + }, # Allows `setup.py test` to work correctly with pytest setup_requires=[] + pytest_runner, # Additional entries you may want simply uncomment the lines you want and fill in the data