{ "cells": [ { "cell_type": "markdown", "id": "8c6f4926", "metadata": {}, "source": [ "# Running `astra benchmark` on a synthetic dataset\n", "\n", "In this tutorial, we show how to run `astra benchmark` on a synthetic classification dataset. We begin by creating the dataset and bringing it into an ASTRA-ready form. ASTRA expects a `pd.DataFrame` with one row per sample, a `Features` column containing the feature vector for each sample, a `Target` column containing the label, and a `Fold` column containing the 0-indexed CV fold assignment. The helper function below creates this format from a standard NumPy feature matrix and label array using stratified 5-fold splits." ] }, { "cell_type": "code", "execution_count": 1, "id": "eb991c7d", "metadata": {}, "outputs": [], "source": [ "import pickle\n", "\n", "import numpy as np\n", "import pandas as pd\n", "from sklearn.datasets import make_classification\n", "from sklearn.metrics import (\n", " average_precision_score,\n", " matthews_corrcoef,\n", " roc_auc_score,\n", ")\n", "from sklearn.model_selection import StratifiedKFold, train_test_split\n" ] }, { "cell_type": "code", "execution_count": 2, "id": "4b19d961", "metadata": {}, "outputs": [], "source": [ "def prepare_for_ASTRA(X_train: np.ndarray, y_train: np.ndarray) -> pd.DataFrame:\n", " \"\"\"\n", " Prepare the training data for ASTRA by creating 5-fold cross-validation splits\n", " and returning a pd.DataFrame with the following columns:\n", " - 'Features': a list of feature vectors for each sample\n", " - 'Target': the corresponding label for each sample\n", " - 'Fold': the fold number (0-4) for cross-validation\n", "\n", " Parameters\n", " ----------\n", " X_train : array-like, shape (n_samples, n_features)\n", " The training feature matrix.\n", " y_train : array-like, shape (n_samples,)\n", " The training labels.\n", "\n", " Returns\n", " -------\n", " pd.DataFrame\n", " A DataFrame containing the features, target labels, and fold assignments.\n", " \"\"\"\n", " kf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)\n", " folds = np.empty(len(X_train), dtype=int)\n", " for fold, (_, val_index) in enumerate(kf.split(X_train, y_train)):\n", " folds[val_index] = fold\n", " df = pd.DataFrame({\"Features\": list(X_train), \"Target\": y_train, \"Fold\": folds})\n", " return df" ] }, { "cell_type": "code", "execution_count": 3, "id": "d3ddc568", "metadata": {}, "outputs": [], "source": [ "X, y = make_classification(\n", " n_samples=1000,\n", " n_features=100,\n", " n_informative=75,\n", " n_redundant=10,\n", " n_repeated=5,\n", " random_state=42,\n", " weights=[0.7, 0.3],\n", ")\n", "\n", "# set aside 10% of the data as test set\n", "X_train, X_test, y_train, y_test = train_test_split(\n", " X, y, test_size=0.1, random_state=42, stratify=y\n", ")\n", "\n", "df_train = prepare_for_ASTRA(X_train, y_train)\n", "df_train.to_pickle(\"synthetic_data_train.pkl\")\n", "\n", "with open(\"synthetic_data_test.pkl\", \"wb\") as f:\n", " pickle.dump((X_test, y_test), f)" ] }, { "cell_type": "markdown", "id": "a6c067ab", "metadata": {}, "source": [ "Now we can run `astra benchmark`. We use ROC-AUC as the main metric for model selection, with PR-AUC and MCC as secondary metrics. ASTRA will evaluate all built-in classifiers 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`)." ] }, { "cell_type": "code", "execution_count": 4, "id": "4e2fc648", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ " \n", " \u001b[1;34m👋 Welcome to ASTRA - Automated model selection using statistical testing\u001b[0m \n", "\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m ----------------------------------------------\u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m _ \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m __ _ ___ | |_ _ __ __ _ \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\u001b[1;35m/\u001b[0m\u001b[1;34m _` |\u001b[0m\u001b[1;35m/\u001b[0m\u001b[1;34m __|| __|| '__| \u001b[0m\u001b[1;35m/\u001b[0m\u001b[1;34m _` | \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m | \u001b[0m\u001b[1;34m(\u001b[0m\u001b[1;34m_| |\\__ \\| |_ | | | \u001b[0m\u001b[1;34m(\u001b[0m\u001b[1;34m_| | \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m \\__,_||___/ \\__||_| \\__,_| \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m ----------------------------------------------\u001b[0m\u001b[1;34m \u001b[0m\n", "\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\u001b[1;34m \u001b[0m\n", " \u001b[1;36m🤔 For help, run: astra --help\u001b[0m \n", " \u001b[1;36m🧪 To benchmark models, run: astra benchmark --help\u001b[0m \n", " \u001b[1;36m🏆 To compare models, run: astra compare --help\u001b[0m \n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Starting benchmark for synthetic_data_benchmark.\n", "15-03 10:17 - INFO: Loading data.\n", "15-03 10:17 - INFO: Starting benchmarking.\n", "15-03 10:17 - INFO: Features column: Features\n", "15-03 10:17 - INFO: Target column: Target\n", "15-03 10:17 - INFO: Running 5-fold CV.\n", "15-03 10:17 - INFO: Fold column: Fold\n", "15-03 10:17 - INFO: Using Optuna for hyperparameter optimization, with 100 trials and a timeout of 3600 seconds.\n", "15-03 10:17 - INFO: Will check assumptions for parametric tests and use them if met.\n", "15-03 10:17 - INFO: Getting models and parameters.\n", "15-03 10:17 - INFO: Benchmarking classification models.\n", "15-03 10:17 - INFO: Main metric: roc_auc\n", "15-03 10:17 - INFO: Secondary metrics: ['pr_auc', 'mcc']\n", "15-03 10:17 - INFO: Starting CV for all models using default hyperparameters.\n", "15-03 10:17 - INFO: Running CV.\n", "15-03 10:17 - INFO: Running LogisticRegression.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for LogisticRegression:\n", " roc_auc: 0.915 ± 0.018 (median: 0.918)\n", " pr_auc: 0.854 ± 0.040 (median: 0.862)\n", " mcc: 0.686 ± 0.056 (median: 0.674)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running GaussianProcessClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for GaussianProcessClassifier:\n", " roc_auc: 0.500 ± 0.000 (median: 0.500)\n", " pr_auc: 0.299 ± 0.002 (median: 0.300)\n", " mcc: 0.000 ± 0.000 (median: 0.000)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running BernoulliNB.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for BernoulliNB:\n", " roc_auc: 0.828 ± 0.032 (median: 0.822)\n", " pr_auc: 0.686 ± 0.051 (median: 0.673)\n", " mcc: 0.473 ± 0.095 (median: 0.421)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running GaussianNB.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for GaussianNB:\n", " roc_auc: 0.903 ± 0.022 (median: 0.893)\n", " pr_auc: 0.823 ± 0.047 (median: 0.797)\n", " mcc: 0.605 ± 0.100 (median: 0.574)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running DecisionTreeClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for DecisionTreeClassifier:\n", " roc_auc: 0.609 ± 0.021 (median: 0.620)\n", " pr_auc: 0.369 ± 0.019 (median: 0.379)\n", " mcc: 0.219 ± 0.046 (median: 0.242)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running ExtraTreeClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for ExtraTreeClassifier:\n", " roc_auc: 0.570 ± 0.038 (median: 0.593)\n", " pr_auc: 0.341 ± 0.024 (median: 0.354)\n", " mcc: 0.141 ± 0.076 (median: 0.180)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running ExtraTreesClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for ExtraTreesClassifier:\n", " roc_auc: 0.918 ± 0.014 (median: 0.913)\n", " pr_auc: 0.852 ± 0.028 (median: 0.852)\n", " mcc: 0.328 ± 0.059 (median: 0.329)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running RandomForestClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for RandomForestClassifier:\n", " roc_auc: 0.896 ± 0.015 (median: 0.896)\n", " pr_auc: 0.805 ± 0.045 (median: 0.823)\n", " mcc: 0.361 ± 0.049 (median: 0.370)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running GradientBoostingClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for GradientBoostingClassifier:\n", " roc_auc: 0.895 ± 0.025 (median: 0.901)\n", " pr_auc: 0.820 ± 0.050 (median: 0.814)\n", " mcc: 0.568 ± 0.053 (median: 0.559)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running BaggingClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for BaggingClassifier:\n", " roc_auc: 0.773 ± 0.049 (median: 0.794)\n", " pr_auc: 0.587 ± 0.068 (median: 0.610)\n", " mcc: 0.351 ± 0.070 (median: 0.311)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running HistGradientBoostingClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for HistGradientBoostingClassifier:\n", " roc_auc: 0.928 ± 0.009 (median: 0.925)\n", " pr_auc: 0.874 ± 0.016 (median: 0.872)\n", " mcc: 0.651 ± 0.069 (median: 0.673)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running AdaBoostClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for AdaBoostClassifier:\n", " roc_auc: 0.826 ± 0.026 (median: 0.828)\n", " pr_auc: 0.691 ± 0.064 (median: 0.674)\n", " mcc: 0.461 ± 0.066 (median: 0.472)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running KNeighborsClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for KNeighborsClassifier:\n", " roc_auc: 0.726 ± 0.038 (median: 0.730)\n", " pr_auc: 0.507 ± 0.031 (median: 0.502)\n", " mcc: 0.293 ± 0.076 (median: 0.309)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running NearestCentroid.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for NearestCentroid:\n", " roc_auc: 0.893 ± 0.029 (median: 0.888)\n", " pr_auc: 0.803 ± 0.055 (median: 0.798)\n", " mcc: 0.184 ± 0.030 (median: 0.200)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running LinearDiscriminantAnalysis.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for LinearDiscriminantAnalysis:\n", " roc_auc: 0.928 ± 0.015 (median: 0.922)\n", " pr_auc: 0.881 ± 0.034 (median: 0.879)\n", " mcc: 0.736 ± 0.061 (median: 0.742)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running SGDClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for SGDClassifier:\n", " roc_auc: 0.822 ± 0.047 (median: 0.828)\n", " pr_auc: 0.660 ± 0.085 (median: 0.688)\n", " mcc: 0.615 ± 0.105 (median: 0.668)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running MLPClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for MLPClassifier:\n", " roc_auc: 0.948 ± 0.015 (median: 0.949)\n", " pr_auc: 0.897 ± 0.046 (median: 0.914)\n", " mcc: 0.738 ± 0.062 (median: 0.712)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running LGBMClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for LGBMClassifier:\n", " roc_auc: 0.933 ± 0.014 (median: 0.924)\n", " pr_auc: 0.885 ± 0.022 (median: 0.885)\n", " mcc: 0.639 ± 0.030 (median: 0.644)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Running XGBClassifier.\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " Performance for XGBClassifier:\n", " roc_auc: 0.918 ± 0.015 (median: 0.915)\n", " pr_auc: 0.857 ± 0.024 (median: 0.851)\n", " mcc: 0.651 ± 0.033 (median: 0.667)\n" ] }, { "name": "stderr", "output_type": "stream", "text": [ "15-03 10:17 - INFO: Done!\n", "15-03 10:17 - INFO: Finished CV for all models.\n", "15-03 10:17 - INFO: Checking assumptions for parametric tests.\n", "15-03 10:17 - INFO: Assumptions of parametric tests met: False.\n", "15-03 10:17 - INFO: Finding best model.\n", "15-03 10:17 - INFO: Best model: MLPClassifier. Reason: Conover post-hoc test.\n", "15-03 10:17 - INFO: Starting final hyperparameter tuning.\n", "15-03 10:18 - INFO: Done!\n" ] }, { "name": "stdout", "output_type": "stream", "text": [ " -------------\n", " Final results\n", " -------------\n", " Final model: MLPClassifier\n", " Hyperparameters:\n", " hidden_layer_sizes: (256, 128, 64)\n", " activation: relu\n", " solver: adam\n", " alpha: 0.00028665421960535865\n", " learning_rate: constant\n", " learning_rate_init: 0.004806645399496772\n", " early_stopping: False\n", " batch_size: 204\n", " Mean roc_auc: 0.963 ± 0.007.\n", " Median roc_auc: 0.965.\n", " Mean pr_auc: 0.934 ± 0.018.\n", " Median pr_auc: 0.929.\n", " Mean mcc: 0.794 ± 0.056.\n", " Median mcc: 0.803.\n" ] } ], "source": [ "%%bash\n", "astra benchmark synthetic_data_train.pkl --name synthetic_data_benchmark --use_optuna --main_metric roc_auc --sec_metrics pr_auc mcc" ] }, { "cell_type": "markdown", "id": "b7bc7564", "metadata": {}, "source": [ "The output tells us that `MLPClassifier` was selected as the best model, which means that it performed significantly better than the other models. Parametric test assumptions were not met, so non-parametric tests were used. The output was saved as `benchmark.log` under `results/synthetic_data_benchmark`. The final model was saved as `final_model.pkl`.\n", "\n", "ASTRA also saved the model's final performance (`final_CV.pkl`) and optimised hyperparameters (`final_hyperparameters.pkl`):" ] }, { "cell_type": "code", "execution_count": 5, "id": "46324a2c", "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "{'roc_auc': '0.963 ± 0.007', 'pr_auc': '0.934 ± 0.018', 'mcc': '0.794 ± 0.056'}\n", "{'hidden_layer_sizes': (256, 128, 64), 'activation': 'relu', 'solver': 'adam', 'alpha': 0.00028665421960535865, 'learning_rate': 'constant', 'learning_rate_init': 0.004806645399496772, 'early_stopping': False, 'batch_size': 204}\n" ] } ], "source": [ "with open(\"results/synthetic_data_benchmark/final_CV.pkl\", \"rb\") as f:\n", " final_CV = pickle.load(f)\n", "print({k: f\"{np.mean(v):.3f} ± {np.std(v):.3f}\" for k, v in final_CV.items()})\n", "\n", "with open(\"results/synthetic_data_benchmark/final_hyperparameters.pkl\", \"rb\") as f:\n", " final_hyperparameters = pickle.load(f)\n", "print(final_hyperparameters)" ] }, { "cell_type": "markdown", "id": "816d0deb", "metadata": {}, "source": [ "We can also see the performance metrics of all models using default hyperparameters, which are saved in `default_CV.pkl`:" ] }, { "cell_type": "code", "execution_count": null, "id": "dddfe22c", "metadata": {}, "outputs": [ { "data": { "text/html": [ "
| \n", " | roc_auc | \n", "pr_auc | \n", "mcc | \n", "
|---|---|---|---|
| MLPClassifier | \n", "0.948 ± 0.015 | \n", "0.897 ± 0.046 | \n", "0.738 ± 0.062 | \n", "
| LGBMClassifier | \n", "0.933 ± 0.014 | \n", "0.885 ± 0.022 | \n", "0.639 ± 0.030 | \n", "
| LinearDiscriminantAnalysis | \n", "0.928 ± 0.015 | \n", "0.881 ± 0.034 | \n", "0.736 ± 0.061 | \n", "
| HistGradientBoostingClassifier | \n", "0.928 ± 0.009 | \n", "0.874 ± 0.016 | \n", "0.651 ± 0.069 | \n", "
| XGBClassifier | \n", "0.918 ± 0.015 | \n", "0.857 ± 0.024 | \n", "0.651 ± 0.033 | \n", "
| ExtraTreesClassifier | \n", "0.918 ± 0.014 | \n", "0.852 ± 0.028 | \n", "0.328 ± 0.059 | \n", "
| LogisticRegression | \n", "0.915 ± 0.018 | \n", "0.854 ± 0.040 | \n", "0.686 ± 0.056 | \n", "
| GaussianNB | \n", "0.903 ± 0.022 | \n", "0.823 ± 0.047 | \n", "0.605 ± 0.100 | \n", "
| RandomForestClassifier | \n", "0.896 ± 0.015 | \n", "0.805 ± 0.045 | \n", "0.361 ± 0.049 | \n", "
| GradientBoostingClassifier | \n", "0.895 ± 0.025 | \n", "0.820 ± 0.050 | \n", "0.568 ± 0.053 | \n", "
| NearestCentroid | \n", "0.893 ± 0.029 | \n", "0.803 ± 0.055 | \n", "0.184 ± 0.030 | \n", "
| BernoulliNB | \n", "0.828 ± 0.032 | \n", "0.686 ± 0.051 | \n", "0.473 ± 0.095 | \n", "
| AdaBoostClassifier | \n", "0.826 ± 0.026 | \n", "0.691 ± 0.064 | \n", "0.461 ± 0.066 | \n", "
| SGDClassifier | \n", "0.822 ± 0.047 | \n", "0.660 ± 0.085 | \n", "0.615 ± 0.105 | \n", "
| BaggingClassifier | \n", "0.773 ± 0.049 | \n", "0.587 ± 0.068 | \n", "0.351 ± 0.070 | \n", "
| KNeighborsClassifier | \n", "0.726 ± 0.038 | \n", "0.507 ± 0.031 | \n", "0.293 ± 0.076 | \n", "
| DecisionTreeClassifier | \n", "0.609 ± 0.021 | \n", "0.369 ± 0.019 | \n", "0.219 ± 0.046 | \n", "
| ExtraTreeClassifier | \n", "0.570 ± 0.038 | \n", "0.341 ± 0.024 | \n", "0.141 ± 0.076 | \n", "
| GaussianProcessClassifier | \n", "0.500 ± 0.000 | \n", "0.299 ± 0.002 | \n", "0.000 ± 0.000 | \n", "