Case study on real-life drug discovery data

In this tutorial, we demonstrate the complete ASTRA workflow on real-life drug discovery data from the ASAP Discovery x OpenADMET Antiviral Challenge. A part of the challenge is predicting the LogD of several hundred chemical compounds. LogD is a measure of how a compound distributes between a lipid phase and water, and important for drug discovery because it affects membrane permeability, solubility, and bioavailability.

The data were split by when they were tested, with compounds from early stages of the drug discovery campaign serving as the training data, and compounds from later stages serving as the test data. The test data were additionally stratified by chemical similarity, removing compounds that were deemed too chemically similar to training compounds. This challenge therefore tests ML models retrospectively in real-life drug discovery settings, where ML models can be trained on early data to guide further drug optimisation efforts.

We split the training data using k-means clustering and calculated six standard cheminformatics fingerprints (see ASTRA's benchmark repository). We provide ASTRA-ready datasets here.

A typical ASTRA workflow consists of:

  1. running astra benchmark for every fingerprint, yielding a single model per fingerprint, and

  2. running astra compare to compare models obtained for different fingerprints.

Running astra benchmark for every fingerprint

We run astra benchmark with MSE as the main metric for model selection, and R2 and MAE as secondary metrics. ASTRA will evaluate all built-in regressors using 5-fold CV and automatically select the best one via statistical testing. Hyperparameter tuning of the final model is performed using Optuna (--use_optuna) with a timeout of 100 s.

%%bash
astra benchmark features/LogD_atompair_train.pkl \
    --name LogD_atompair \
    --use_optuna \
    --fold_col KMeans_Cluster_42 \
    --main_metric MSE \
    --sec_metrics R2 MAE \
    --timeout 100
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 14:55 - INFO: Starting benchmark for LogD_atompair.
15-03 14:55 - INFO: Loading data.
15-03 14:55 - INFO: Starting benchmarking.
15-03 14:55 - INFO: Features column: Features
15-03 14:55 - INFO: Target column: Target
15-03 14:55 - INFO: Running 5-fold CV.
15-03 14:55 - INFO: Fold column: KMeans_Cluster_42
15-03 14:55 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 100 seconds.
15-03 14:55 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 14:55 - INFO: Getting models and parameters.
15-03 14:55 - INFO: Benchmarking regression models.
15-03 14:55 - INFO: Main metric: mse
15-03 14:55 - INFO: Secondary metrics: ['r2', 'mae']
15-03 14:55 - INFO: Starting CV for all models using default hyperparameters.
15-03 14:55 - INFO: Running CV.
15-03 14:55 - INFO: Running XGBRegressor.
                    Performance for XGBRegressor:
                    mse: 0.956 ± 0.278 (median: 0.792)
                    r2: 0.248 ± 0.418 (median: 0.432)
                    mae: 0.737 ± 0.121 (median: 0.693)
15-03 14:55 - INFO: Running RandomForestRegressor.
                    Performance for RandomForestRegressor:
                    mse: 0.815 ± 0.125 (median: 0.834)
                    r2: 0.401 ± 0.153 (median: 0.467)
                    mae: 0.673 ± 0.076 (median: 0.694)
15-03 14:56 - INFO: Running GradientBoostingRegressor.
                    Performance for GradientBoostingRegressor:
                    mse: 0.856 ± 0.151 (median: 0.908)
                    r2: 0.371 ± 0.175 (median: 0.419)
                    mae: 0.705 ± 0.080 (median: 0.702)
15-03 14:56 - INFO: Running HistGradientBoostingRegressor.
                    Performance for HistGradientBoostingRegressor:
                    mse: 0.857 ± 0.129 (median: 0.875)
                    r2: 0.354 ± 0.235 (median: 0.472)
                    mae: 0.708 ± 0.082 (median: 0.728)
15-03 14:56 - INFO: Running KNeighborsRegressor.
                    Performance for KNeighborsRegressor:
                    mse: 0.932 ± 0.263 (median: 0.820)
                    r2: 0.321 ± 0.180 (median: 0.207)
                    mae: 0.713 ± 0.088 (median: 0.683)
15-03 14:56 - INFO: Running SVR.
                    Performance for SVR:
                    mse: 0.762 ± 0.192 (median: 0.634)
                    r2: 0.456 ± 0.096 (median: 0.450)
                    mae: 0.664 ± 0.096 (median: 0.665)
15-03 14:56 - INFO: Running Ridge.
                    Performance for Ridge:
                    mse: 0.762 ± 0.105 (median: 0.725)
                    r2: 0.419 ± 0.213 (median: 0.523)
                    mae: 0.669 ± 0.054 (median: 0.686)
15-03 14:56 - INFO: Running BayesianRidge.
                    Performance for BayesianRidge:
                    mse: 0.709 ± 0.097 (median: 0.729)
                    r2: 0.470 ± 0.163 (median: 0.537)
                    mae: 0.644 ± 0.052 (median: 0.648)
15-03 14:56 - INFO: Running KernelRidge.
                    Performance for KernelRidge:
                    mse: 0.793 ± 0.161 (median: 0.739)
                    r2: 0.399 ± 0.235 (median: 0.457)
                    mae: 0.673 ± 0.071 (median: 0.651)
15-03 14:56 - INFO: Running LGBMRegressor.
                    Performance for LGBMRegressor:
                    mse: 0.857 ± 0.129 (median: 0.875)
                    r2: 0.354 ± 0.235 (median: 0.472)
                    mae: 0.708 ± 0.082 (median: 0.728)
15-03 14:56 - INFO: Running CatBoostRegressor.
                    Performance for CatBoostRegressor:
                    mse: 0.713 ± 0.138 (median: 0.645)
                    r2: 0.480 ± 0.115 (median: 0.480)
                    mae: 0.639 ± 0.086 (median: 0.658)
15-03 14:57 - INFO: Done!
15-03 14:57 - INFO: Finished CV for all models.
15-03 14:57 - INFO: Checking assumptions for parametric tests.
15-03 14:57 - INFO: Assumptions of parametric tests met: False.
15-03 14:57 - INFO: Finding best model.
15-03 14:57 - INFO: Best model: SVR. Reason: median score.
15-03 14:57 - INFO: Starting final hyperparameter tuning.
15-03 14:57 - INFO: Done!
                    -------------
                    Final results
                    -------------
                    Final model: SVR
                    Hyperparameters:
                    kernel: sigmoid
                    C: 7.867745381977843
                    gamma: 0.001
                    Mean mse: 0.683 ± 0.158.
                    Median mse: 0.628.
                    Mean r2: 0.510 ± 0.100.
                    Median r2: 0.517.
                    Mean mae: 0.633 ± 0.085.
                    Median mae: 0.636.
%%bash
astra benchmark features/LogD_cats2d_train.pkl \
    --name LogD_cats2d \
    --use_optuna \
    --fold_col KMeans_Cluster_42 \
    --main_metric MSE \
    --sec_metrics R2 MAE \
    --timeout 100
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 14:57 - INFO: Starting benchmark for LogD_cats2d.
15-03 14:57 - INFO: Loading data.
15-03 14:57 - INFO: Starting benchmarking.
15-03 14:57 - INFO: Features column: Features
15-03 14:57 - INFO: Target column: Target
15-03 14:57 - INFO: Running 5-fold CV.
15-03 14:57 - INFO: Fold column: KMeans_Cluster_42
15-03 14:57 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 100 seconds.
15-03 14:57 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 14:57 - INFO: Getting models and parameters.
15-03 14:57 - INFO: Benchmarking regression models.
15-03 14:57 - INFO: Main metric: mse
15-03 14:57 - INFO: Secondary metrics: ['r2', 'mae']
15-03 14:57 - INFO: Starting CV for all models using default hyperparameters.
15-03 14:57 - INFO: Running CV.
15-03 14:57 - INFO: Running XGBRegressor.
                    Performance for XGBRegressor:
                    mse: 1.110 ± 0.375 (median: 0.909)
                    r2: 0.224 ± 0.153 (median: 0.240)
                    mae: 0.779 ± 0.125 (median: 0.719)
15-03 14:57 - INFO: Running RandomForestRegressor.
                    Performance for RandomForestRegressor:
                    mse: 0.967 ± 0.258 (median: 0.873)
                    r2: 0.312 ± 0.126 (median: 0.353)
                    mae: 0.757 ± 0.089 (median: 0.732)
15-03 14:57 - INFO: Running GradientBoostingRegressor.
                    Performance for GradientBoostingRegressor:
                    mse: 0.929 ± 0.220 (median: 1.025)
                    r2: 0.341 ± 0.117 (median: 0.382)
                    mae: 0.754 ± 0.114 (median: 0.781)
15-03 14:57 - INFO: Running HistGradientBoostingRegressor.
                    Performance for HistGradientBoostingRegressor:
                    mse: 0.983 ± 0.308 (median: 0.843)
                    r2: 0.309 ± 0.124 (median: 0.309)
                    mae: 0.769 ± 0.109 (median: 0.702)
15-03 14:57 - INFO: Running KNeighborsRegressor.
                    Performance for KNeighborsRegressor:
                    mse: 1.015 ± 0.322 (median: 0.938)
                    r2: 0.293 ± 0.109 (median: 0.287)
                    mae: 0.748 ± 0.111 (median: 0.716)
15-03 14:57 - INFO: Running SVR.
                    Performance for SVR:
                    mse: 1.048 ± 0.133 (median: 1.001)
                    r2: 0.228 ± 0.181 (median: 0.308)
                    mae: 0.822 ± 0.051 (median: 0.812)
15-03 14:57 - INFO: Running Ridge.
                    Performance for Ridge:
                    mse: 1.387 ± 0.327 (median: 1.269)
                    r2: -0.054 ± 0.474 (median: 0.113)
                    mae: 0.887 ± 0.095 (median: 0.887)
15-03 14:57 - INFO: Running BayesianRidge.
                    Performance for BayesianRidge:
                    mse: 0.912 ± 0.208 (median: 0.850)
                    r2: 0.328 ± 0.193 (median: 0.288)
                    mae: 0.753 ± 0.073 (median: 0.714)
15-03 14:57 - INFO: Running KernelRidge.
                    Performance for KernelRidge:
                    mse: 1.395 ± 0.324 (median: 1.283)
                    r2: -0.060 ± 0.474 (median: 0.109)
                    mae: 0.889 ± 0.094 (median: 0.894)
15-03 14:57 - INFO: Running LGBMRegressor.
                    Performance for LGBMRegressor:
                    mse: 0.988 ± 0.324 (median: 0.837)
                    r2: 0.311 ± 0.119 (median: 0.331)
                    mae: 0.771 ± 0.124 (median: 0.692)
15-03 14:57 - INFO: Running CatBoostRegressor.
                    Performance for CatBoostRegressor:
                    mse: 0.845 ± 0.142 (median: 0.914)
                    r2: 0.394 ± 0.070 (median: 0.435)
                    mae: 0.723 ± 0.052 (median: 0.739)
15-03 14:57 - INFO: Done!
15-03 14:57 - INFO: Finished CV for all models.
15-03 14:57 - INFO: Checking assumptions for parametric tests.
15-03 14:57 - INFO: Assumptions of parametric tests met: False.
15-03 14:57 - INFO: Finding best model.
15-03 14:57 - INFO: Best model: CatBoostRegressor. Reason: Conover post-hoc test.
15-03 14:57 - INFO: Starting final hyperparameter tuning.
15-03 14:59 - INFO: Done!
                    -------------
                    Final results
                    -------------
                    Final model: CatBoostRegressor
                    Hyperparameters:
                    iterations: 556
                    learning_rate: 0.016511657063497182
                    depth: 4
                    l2_leaf_reg: 2.235048729002267
                    rsm: 0.3668787683669737
                    loss_function: RMSE
                    border_count: 134
                    feature_border_type: Median
                    random_strength: 1.3063759368181672e-06
                    bootstrap_type: Bayesian
                    Mean mse: 0.887 ± 0.156.
                    Median mse: 0.906.
                    Mean r2: 0.361 ± 0.102.
                    Median r2: 0.360.
                    Mean mae: 0.744 ± 0.053.
                    Median mae: 0.755.
%%bash
astra benchmark features/LogD_desc2D_train.pkl \
    --name LogD_desc2D \
    --use_optuna \
    --fold_col KMeans_Cluster_42 \
    --main_metric MSE \
    --sec_metrics R2 MAE \
    --timeout 100
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 14:59 - INFO: Starting benchmark for LogD_desc2D.
15-03 14:59 - INFO: Loading data.
15-03 14:59 - INFO: Starting benchmarking.
15-03 14:59 - INFO: Features column: Features
15-03 14:59 - INFO: Target column: Target
15-03 14:59 - INFO: Running 5-fold CV.
15-03 14:59 - INFO: Fold column: KMeans_Cluster_42
15-03 14:59 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 100 seconds.
15-03 14:59 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 14:59 - INFO: Getting models and parameters.
15-03 14:59 - INFO: Benchmarking regression models.
15-03 14:59 - INFO: Main metric: mse
15-03 14:59 - INFO: Secondary metrics: ['r2', 'mae']
15-03 14:59 - INFO: Starting CV for all models using default hyperparameters.
15-03 14:59 - INFO: Running CV.
15-03 14:59 - INFO: Running XGBRegressor.
                    Performance for XGBRegressor:
                    mse: 0.806 ± 0.158 (median: 0.839)
                    r2: 0.409 ± 0.170 (median: 0.491)
                    mae: 0.699 ± 0.059 (median: 0.739)
15-03 14:59 - INFO: Running RandomForestRegressor.
                    Performance for RandomForestRegressor:
                    mse: 0.775 ± 0.232 (median: 0.723)
                    r2: 0.452 ± 0.135 (median: 0.438)
                    mae: 0.676 ± 0.130 (median: 0.666)
15-03 14:59 - INFO: Running GradientBoostingRegressor.
                    Performance for GradientBoostingRegressor:
                    mse: 0.652 ± 0.179 (median: 0.690)
                    r2: 0.528 ± 0.148 (median: 0.598)
                    mae: 0.630 ± 0.088 (median: 0.643)
15-03 14:59 - INFO: Running HistGradientBoostingRegressor.
                    Performance for HistGradientBoostingRegressor:
                    mse: 0.609 ± 0.165 (median: 0.603)
                    r2: 0.564 ± 0.117 (median: 0.636)
                    mae: 0.598 ± 0.087 (median: 0.620)
15-03 14:59 - INFO: Running KNeighborsRegressor.
                    Performance for KNeighborsRegressor:
                    mse: 1.244 ± 0.296 (median: 1.266)
                    r2: 0.101 ± 0.226 (median: 0.253)
                    mae: 0.876 ± 0.122 (median: 0.895)
15-03 14:59 - INFO: Running SVR.
                    Performance for SVR:
                    mse: 1.423 ± 0.162 (median: 1.409)
                    r2: -0.049 ± 0.250 (median: 0.076)
                    mae: 0.972 ± 0.041 (median: 0.985)
15-03 14:59 - INFO: Running Ridge.
                    Performance for Ridge:
                    mse: 0.774 ± 0.107 (median: 0.790)
                    r2: 0.425 ± 0.157 (median: 0.441)
                    mae: 0.676 ± 0.044 (median: 0.673)
15-03 14:59 - INFO: Running BayesianRidge.
                    Performance for BayesianRidge:
                    mse: 0.677 ± 0.201 (median: 0.749)
                    r2: 0.506 ± 0.177 (median: 0.548)
                    mae: 0.632 ± 0.087 (median: 0.679)
15-03 14:59 - INFO: Running KernelRidge.
                    Performance for KernelRidge:
                    mse: 0.774 ± 0.107 (median: 0.790)
                    r2: 0.425 ± 0.157 (median: 0.441)
                    mae: 0.676 ± 0.044 (median: 0.673)
15-03 14:59 - INFO: Running LGBMRegressor.
                    Performance for LGBMRegressor:
                    mse: 0.612 ± 0.207 (median: 0.598)
                    r2: 0.566 ± 0.135 (median: 0.639)
                    mae: 0.594 ± 0.099 (median: 0.606)
15-03 14:59 - INFO: Running CatBoostRegressor.
                    Performance for CatBoostRegressor:
                    mse: 0.653 ± 0.145 (median: 0.681)
                    r2: 0.527 ± 0.128 (median: 0.589)
                    mae: 0.622 ± 0.091 (median: 0.647)
15-03 14:59 - INFO: Done!
15-03 14:59 - INFO: Finished CV for all models.
15-03 14:59 - INFO: Checking assumptions for parametric tests.
15-03 14:59 - INFO: Assumptions of parametric tests met: False.
15-03 14:59 - INFO: Finding best model.
15-03 14:59 - INFO: Best model: LGBMRegressor. Reason: Conover post-hoc test.
15-03 14:59 - INFO: Starting final hyperparameter tuning.
15-03 15:01 - INFO: Done!
                    -------------
                    Final results
                    -------------
                    Final model: LGBMRegressor
                    Hyperparameters:
                    boosting_type: dart
                    num_leaves: 510
                    max_depth: 11
                    learning_rate: 0.15869372433797996
                    n_estimators: 563
                    min_data_in_leaf: 23
                    subsample: 0.6769996225690526
                    subsample_freq: 4
                    colsample_bytree: 0.7557762135754856
                    reg_alpha: 0.6705738213879265
                    reg_lambda: 0.5707565631892096
                    bagging_freq: 10
                    min_split_gain: 0.03202064627142526
                    min_child_weight: 8.23573947779089
                    min_child_samples: 5
                    Mean mse: 0.609 ± 0.154.
                    Median mse: 0.624.
                    Mean r2: 0.563 ± 0.107.
                    Median r2: 0.613.
                    Mean mae: 0.599 ± 0.073.
                    Median mae: 0.610.
%%bash
astra benchmark features/LogD_ecfp_train.pkl \
    --name LogD_ecfp \
    --use_optuna \
    --fold_col KMeans_Cluster_42 \
    --main_metric MSE \
    --sec_metrics R2 MAE \
    --timeout 100
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 15:01 - INFO: Starting benchmark for LogD_ecfp.
15-03 15:01 - INFO: Loading data.
15-03 15:01 - INFO: Starting benchmarking.
15-03 15:01 - INFO: Features column: Features
15-03 15:01 - INFO: Target column: Target
15-03 15:01 - INFO: Running 5-fold CV.
15-03 15:01 - INFO: Fold column: KMeans_Cluster_42
15-03 15:01 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 100 seconds.
15-03 15:01 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 15:01 - INFO: Getting models and parameters.
15-03 15:01 - INFO: Benchmarking regression models.
15-03 15:01 - INFO: Main metric: mse
15-03 15:01 - INFO: Secondary metrics: ['r2', 'mae']
15-03 15:01 - INFO: Starting CV for all models using default hyperparameters.
15-03 15:01 - INFO: Running CV.
15-03 15:01 - INFO: Running XGBRegressor.
                    Performance for XGBRegressor:
                    mse: 1.003 ± 0.133 (median: 0.995)
                    r2: 0.260 ± 0.172 (median: 0.278)
                    mae: 0.760 ± 0.061 (median: 0.734)
15-03 15:01 - INFO: Running RandomForestRegressor.
                    Performance for RandomForestRegressor:
                    mse: 0.899 ± 0.070 (median: 0.878)
                    r2: 0.323 ± 0.197 (median: 0.403)
                    mae: 0.725 ± 0.031 (median: 0.723)
15-03 15:01 - INFO: Running GradientBoostingRegressor.
                    Performance for GradientBoostingRegressor:
                    mse: 1.039 ± 0.144 (median: 1.143)
                    r2: 0.229 ± 0.212 (median: 0.330)
                    mae: 0.778 ± 0.054 (median: 0.792)
15-03 15:01 - INFO: Running HistGradientBoostingRegressor.
                    Performance for HistGradientBoostingRegressor:
                    mse: 0.832 ± 0.079 (median: 0.795)
                    r2: 0.385 ± 0.140 (median: 0.431)
                    mae: 0.708 ± 0.047 (median: 0.694)
15-03 15:01 - INFO: Running KNeighborsRegressor.
                    Performance for KNeighborsRegressor:
                    mse: 0.816 ± 0.091 (median: 0.795)
                    r2: 0.402 ± 0.120 (median: 0.468)
                    mae: 0.681 ± 0.048 (median: 0.653)
15-03 15:01 - INFO: Running SVR.
                    Performance for SVR:
                    mse: 0.839 ± 0.160 (median: 0.752)
                    r2: 0.392 ± 0.114 (median: 0.418)
                    mae: 0.711 ± 0.077 (median: 0.677)
15-03 15:01 - INFO: Running Ridge.
                    Performance for Ridge:
                    mse: 0.906 ± 0.260 (median: 0.836)
                    r2: 0.343 ± 0.177 (median: 0.352)
                    mae: 0.704 ± 0.091 (median: 0.678)
15-03 15:01 - INFO: Running BayesianRidge.
                    Performance for BayesianRidge:
                    mse: 0.852 ± 0.237 (median: 0.722)
                    r2: 0.386 ± 0.143 (median: 0.362)
                    mae: 0.690 ± 0.079 (median: 0.648)
15-03 15:01 - INFO: Running KernelRidge.
                    Performance for KernelRidge:
                    mse: 0.939 ± 0.279 (median: 0.898)
                    r2: 0.320 ± 0.194 (median: 0.370)
                    mae: 0.709 ± 0.098 (median: 0.703)
15-03 15:01 - INFO: Running LGBMRegressor.
                    Performance for LGBMRegressor:
                    mse: 0.832 ± 0.079 (median: 0.795)
                    r2: 0.385 ± 0.140 (median: 0.431)
                    mae: 0.708 ± 0.047 (median: 0.694)
15-03 15:01 - INFO: Running CatBoostRegressor.
                    Performance for CatBoostRegressor:
                    mse: 0.938 ± 0.208 (median: 0.977)
                    r2: 0.319 ± 0.164 (median: 0.291)
                    mae: 0.733 ± 0.090 (median: 0.729)
15-03 15:01 - INFO: Done!
15-03 15:01 - INFO: Finished CV for all models.
15-03 15:01 - INFO: Checking assumptions for parametric tests.
15-03 15:01 - INFO: Assumptions of parametric tests met: True.
15-03 15:01 - INFO: Finding best model.
15-03 15:01 - INFO: Best model: KNeighborsRegressor. Reason: paired t-test.
15-03 15:01 - INFO: Starting final hyperparameter tuning.
15-03 15:01 - INFO: Done!
                    -------------
                    Final results
                    -------------
                    Final model: KNeighborsRegressor
                    Hyperparameters:
                    n_neighbors: 6
                    weights: distance
                    p: 1
                    Mean mse: 0.790 ± 0.095.
                    Median mse: 0.788.
                    Mean r2: 0.419 ± 0.126.
                    Median r2: 0.481.
                    Mean mae: 0.664 ± 0.051.
                    Median mae: 0.639.
%%bash
astra benchmark features/LogD_maccs_train.pkl \
    --name LogD_maccs \
    --use_optuna \
    --fold_col KMeans_Cluster_42 \
    --main_metric MSE \
    --sec_metrics R2 MAE \
    --timeout 100
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 15:02 - INFO: Starting benchmark for LogD_maccs.
15-03 15:02 - INFO: Loading data.
15-03 15:02 - INFO: Starting benchmarking.
15-03 15:02 - INFO: Features column: Features
15-03 15:02 - INFO: Target column: Target
15-03 15:02 - INFO: Running 5-fold CV.
15-03 15:02 - INFO: Fold column: KMeans_Cluster_42
15-03 15:02 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 100 seconds.
15-03 15:02 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 15:02 - INFO: Getting models and parameters.
15-03 15:02 - INFO: Benchmarking regression models.
15-03 15:02 - INFO: Main metric: mse
15-03 15:02 - INFO: Secondary metrics: ['r2', 'mae']
15-03 15:02 - INFO: Starting CV for all models using default hyperparameters.
15-03 15:02 - INFO: Running CV.
15-03 15:02 - INFO: Running XGBRegressor.
                    Performance for XGBRegressor:
                    mse: 1.443 ± 0.132 (median: 1.412)
                    r2: -0.096 ± 0.377 (median: 0.094)
                    mae: 0.913 ± 0.060 (median: 0.916)
15-03 15:02 - INFO: Running RandomForestRegressor.
                    Performance for RandomForestRegressor:
                    mse: 1.076 ± 0.069 (median: 1.060)
                    r2: 0.192 ± 0.237 (median: 0.365)
                    mae: 0.779 ± 0.054 (median: 0.749)
15-03 15:02 - INFO: Running GradientBoostingRegressor.
                    Performance for GradientBoostingRegressor:
                    mse: 1.068 ± 0.219 (median: 1.157)
                    r2: 0.201 ± 0.307 (median: 0.337)
                    mae: 0.766 ± 0.095 (median: 0.758)
15-03 15:02 - INFO: Running HistGradientBoostingRegressor.
                    Performance for HistGradientBoostingRegressor:
                    mse: 0.974 ± 0.159 (median: 1.032)
                    r2: 0.299 ± 0.103 (median: 0.347)
                    mae: 0.736 ± 0.069 (median: 0.750)
15-03 15:02 - INFO: Running KNeighborsRegressor.
                    Performance for KNeighborsRegressor:
                    mse: 1.223 ± 0.425 (median: 1.044)
                    r2: 0.133 ± 0.212 (median: 0.191)
                    mae: 0.783 ± 0.128 (median: 0.751)
15-03 15:02 - INFO: Running SVR.
                    Performance for SVR:
                    mse: 0.890 ± 0.180 (median: 0.866)
                    r2: 0.357 ± 0.125 (median: 0.356)
                    mae: 0.734 ± 0.063 (median: 0.723)
15-03 15:02 - INFO: Running Ridge.
                    Performance for Ridge:
                    mse: 1.270 ± 0.389 (median: 1.386)
                    r2: 0.072 ± 0.349 (median: 0.164)
                    mae: 0.833 ± 0.134 (median: 0.904)
15-03 15:02 - INFO: Running BayesianRidge.
                    Performance for BayesianRidge:
                    mse: 0.984 ± 0.271 (median: 0.925)
                    r2: 0.300 ± 0.144 (median: 0.370)
                    mae: 0.757 ± 0.079 (median: 0.739)
15-03 15:02 - INFO: Running KernelRidge.
                    Performance for KernelRidge:
                    mse: 1.265 ± 0.389 (median: 1.391)
                    r2: 0.073 ± 0.365 (median: 0.161)
                    mae: 0.833 ± 0.138 (median: 0.911)
15-03 15:02 - INFO: Running LGBMRegressor.
                    Performance for LGBMRegressor:
                    mse: 0.974 ± 0.159 (median: 1.032)
                    r2: 0.299 ± 0.103 (median: 0.347)
                    mae: 0.736 ± 0.069 (median: 0.750)
15-03 15:02 - INFO: Running CatBoostRegressor.
                    Performance for CatBoostRegressor:
                    mse: 0.932 ± 0.155 (median: 0.844)
                    r2: 0.300 ± 0.242 (median: 0.372)
                    mae: 0.735 ± 0.071 (median: 0.718)
15-03 15:02 - INFO: Done!
15-03 15:02 - INFO: Finished CV for all models.
15-03 15:02 - INFO: Checking assumptions for parametric tests.
15-03 15:02 - INFO: Assumptions of parametric tests met: False.
15-03 15:02 - INFO: Finding best model.
15-03 15:02 - INFO: Best model: CatBoostRegressor. Reason: Conover post-hoc test.
15-03 15:02 - INFO: Starting final hyperparameter tuning.
15-03 15:04 - INFO: Done!
                    -------------
                    Final results
                    -------------
                    Final model: CatBoostRegressor
                    Hyperparameters:
                    iterations: 637
                    learning_rate: 0.07140177292623585
                    depth: 2
                    l2_leaf_reg: 0.10030457398605436
                    rsm: 0.5295608684611183
                    loss_function: RMSE
                    border_count: 33
                    feature_border_type: Median
                    random_strength: 0.0012590108285256113
                    bootstrap_type: Bayesian
                    Mean mse: 0.938 ± 0.221.
                    Median mse: 0.911.
                    Mean r2: 0.288 ± 0.307.
                    Median r2: 0.368.
                    Mean mae: 0.731 ± 0.087.
                    Median mae: 0.698.
%%bash
astra benchmark features/LogD_topological_train.pkl \
    --name LogD_topological \
    --use_optuna \
    --fold_col KMeans_Cluster_42 \
    --main_metric MSE \
    --sec_metrics R2 MAE \
    --timeout 100
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 15:04 - INFO: Starting benchmark for LogD_topological.
15-03 15:04 - INFO: Loading data.
15-03 15:04 - INFO: Starting benchmarking.
15-03 15:04 - INFO: Features column: Features
15-03 15:04 - INFO: Target column: Target
15-03 15:04 - INFO: Running 5-fold CV.
15-03 15:04 - INFO: Fold column: KMeans_Cluster_42
15-03 15:04 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 100 seconds.
15-03 15:04 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 15:04 - INFO: Getting models and parameters.
15-03 15:04 - INFO: Benchmarking regression models.
15-03 15:04 - INFO: Main metric: mse
15-03 15:04 - INFO: Secondary metrics: ['r2', 'mae']
15-03 15:04 - INFO: Starting CV for all models using default hyperparameters.
15-03 15:04 - INFO: Running CV.
15-03 15:04 - INFO: Running XGBRegressor.
                    Performance for XGBRegressor:
                    mse: 0.917 ± 0.227 (median: 0.990)
                    r2: 0.320 ± 0.256 (median: 0.422)
                    mae: 0.730 ± 0.095 (median: 0.756)
15-03 15:04 - INFO: Running RandomForestRegressor.
                    Performance for RandomForestRegressor:
                    mse: 0.920 ± 0.191 (median: 0.885)
                    r2: 0.290 ± 0.326 (median: 0.465)
                    mae: 0.728 ± 0.076 (median: 0.694)
15-03 15:04 - INFO: Running GradientBoostingRegressor.
                    Performance for GradientBoostingRegressor:
                    mse: 0.968 ± 0.136 (median: 1.045)
                    r2: 0.275 ± 0.240 (median: 0.390)
                    mae: 0.758 ± 0.063 (median: 0.763)
15-03 15:04 - INFO: Running HistGradientBoostingRegressor.
                    Performance for HistGradientBoostingRegressor:
                    mse: 0.824 ± 0.128 (median: 0.788)
                    r2: 0.397 ± 0.127 (median: 0.414)
                    mae: 0.694 ± 0.086 (median: 0.684)
15-03 15:04 - INFO: Running KNeighborsRegressor.
                    Performance for KNeighborsRegressor:
                    mse: 0.825 ± 0.106 (median: 0.860)
                    r2: 0.397 ± 0.125 (median: 0.465)
                    mae: 0.679 ± 0.045 (median: 0.682)
15-03 15:04 - INFO: Running SVR.
                    Performance for SVR:
                    mse: 0.823 ± 0.162 (median: 0.764)
                    r2: 0.401 ± 0.128 (median: 0.404)
                    mae: 0.698 ± 0.078 (median: 0.683)
15-03 15:04 - INFO: Running Ridge.
                    Performance for Ridge:
                    mse: 0.929 ± 0.181 (median: 0.953)
                    r2: 0.319 ± 0.182 (median: 0.366)
                    mae: 0.726 ± 0.065 (median: 0.711)
15-03 15:04 - INFO: Running BayesianRidge.
                    Performance for BayesianRidge:
                    mse: 0.853 ± 0.167 (median: 0.791)
                    r2: 0.380 ± 0.135 (median: 0.382)
                    mae: 0.704 ± 0.069 (median: 0.668)
15-03 15:04 - INFO: Running KernelRidge.
                    Performance for KernelRidge:
                    mse: 0.925 ± 0.183 (median: 0.966)
                    r2: 0.320 ± 0.193 (median: 0.370)
                    mae: 0.725 ± 0.065 (median: 0.720)
15-03 15:04 - INFO: Running LGBMRegressor.
                    Performance for LGBMRegressor:
                    mse: 0.824 ± 0.128 (median: 0.788)
                    r2: 0.397 ± 0.127 (median: 0.414)
                    mae: 0.694 ± 0.086 (median: 0.684)
15-03 15:04 - INFO: Running CatBoostRegressor.
                    Performance for CatBoostRegressor:
                    mse: 0.858 ± 0.116 (median: 0.853)
                    r2: 0.366 ± 0.161 (median: 0.443)
                    mae: 0.710 ± 0.055 (median: 0.735)
15-03 15:04 - INFO: Done!
15-03 15:04 - INFO: Finished CV for all models.
15-03 15:04 - INFO: Checking assumptions for parametric tests.
15-03 15:04 - INFO: Assumptions of parametric tests met: False.
15-03 15:04 - INFO: Finding best model.
15-03 15:04 - INFO: Best model: SVR. Reason: median score.
15-03 15:04 - INFO: Starting final hyperparameter tuning.
15-03 15:05 - INFO: Done!
                    -------------
                    Final results
                    -------------
                    Final model: SVR
                    Hyperparameters:
                    kernel: linear
                    C: 0.035459424622562255
                    gamma: scale
                    Mean mse: 0.778 ± 0.144.
                    Median mse: 0.675.
                    Mean r2: 0.436 ± 0.104.
                    Median r2: 0.445.
                    Mean mae: 0.677 ± 0.065.
                    Median mae: 0.652.

Running astra compare to compare models obtained for different fingerprints

Now we can run astra compare on the results. Again, we use MSE as the main metric, and R2 and MAE as secondary metrics. ASTRA will compare the models trained on different fingerprints and select the best one via statistical testing.

%%bash
astra compare results/LogD_atompair results/LogD_cats2d results/LogD_desc2D results/LogD_ecfp results/LogD_maccs results/LogD_topological --main_metric MSE --sec_metrics R2 MAE
                                                                                
   👋 Welcome to ASTRA - Automated model selection using statistical testing    
                                                                                
                   ----------------------------------------------               
                                         _                                      
                             __ _  ___ | |_  _ __   __ _                        
                            / _` |/ __|| __|| '__| / _` |                       
                           | (_| |\__ \| |_ | |   | (_| |                       
                            \__,_||___/ \__||_|    \__,_|                       
                                                                                
                   ----------------------------------------------               
                                                                                
                         🤔 For help, run: astra --help                         
              🧪 To benchmark models, run: astra benchmark --help               
                🏆 To compare models, run: astra compare --help                 
15-03 15:05 - INFO: Starting comparison of CV results.
15-03 15:05 - INFO: Will check assumptions for parametric tests and use them if met.
15-03 15:05 - INFO: 6 CV results found.
15-03 15:05 - INFO: Checking assumptions for parametric tests.
15-03 15:05 - INFO: Assumptions of parametric tests met: True.
15-03 15:05 - INFO: Best models based on mse:
                    LogD_atompair
                    LogD_cats2d
                    LogD_desc2D
                    LogD_ecfp
                    LogD_maccs
                    LogD_topological
15-03 15:05 - INFO: Best model overall: LogD_desc2D. Reason: Tukey's HSD test.
--------------------------------------------------
Results:
--------------------------------------------------
Mean mse: 0.609 ± 0.154.
Median mse: 0.624.
Mean r2: 0.563 ± 0.107.
 Median r2: 0.613.
Mean mae: 0.599 ± 0.073.
 Median mae: 0.610.
--------------------------------------------------

The output tells us that the best model is LogD_desc2D, the model trained on desc2D fingerprints, based on Tukey's HSD test.

Evaluation on test data

Finally we can evaluate the final model on the test set:

import pickle

import pandas as pd
from sklearn.metrics import (
    mean_absolute_error,
    mean_squared_error,
    r2_score,
)
with open("results/LogD_desc2D/final_model.pkl", "rb") as f:
    final_model = pickle.load(f)
test_data = pd.read_pickle("features/LogD_desc2D_test.pkl")
X_test = pd.DataFrame(test_data["Features"].to_list())
y_test = test_data["Target"].values

predictions = final_model.predict(X_test)
mse = mean_squared_error(y_test, predictions)
r2 = r2_score(y_test, predictions)
mae = mean_absolute_error(y_test, predictions)
print(f"Test MSE: {mse:.4f}")
print(f"Test R2: {r2:.4f}")
print(f"Test MAE: {mae:.4f}")
Test MSE: 0.3143
Test R2: 0.6151
Test MAE: 0.4422

Conclusion

In this tutorial we demonstrated the complete ASTRA workflow on real drug discovery data from the ASAP Discovery x OpenADMET Antiviral Challenge by:

  1. running astra benchmark for each of six cheminformatics fingerprints to select the best model per fingerprint, and

  2. running astra compare to identify the best fingerprint overall.

ASTRA selected desc2D with an LGBMRegressor as the best combination, achieving a test MSE of 0.314, R² of 0.615, and MAE of 0.442 on held-out compounds from later stages of the drug discovery campaign.

For more detail on the available options, see the User Guide.