{ "cells": [ { "cell_type": "markdown", "id": "50f0c7bc", "metadata": {}, "source": [ "# Switching the classification metric with `partial_fit`\n", "\n", "Brush uses the `scorer` to evaluate programs: it drives selection, survival,\n", "the archive, and the final model choice. For classification, the available\n", "scorers are:\n", "\n", "| scorer | binary | multiclass |\n", "|---|---|---|\n", "| `\"log\"` / `\"multi_log\"` | log loss | multinomial log loss |\n", "| `\"accuracy\"` | accuracy | accuracy |\n", "| `\"balanced_accuracy\"` | balanced accuracy | balanced accuracy |\n", "| `\"precision\"` | precision (threshold 0.5) | macro precision |\n", "| `\"recall\"` | recall (threshold 0.5) | macro recall |\n", "| `\"roc_auc\"` | AUROC | macro one-vs-rest AUROC |\n", "| `\"average_precision_score\"` | average precision (AUPRC) | macro one-vs-rest average precision |\n", "\n", "The scorer is **not** used to fit parameters. Weights are always optimized with\n", "the log loss, and split thresholds with the gini impurity. This keeps parameter\n", "fitting smooth and well-behaved while letting you pick whichever metric you care\n", "about for model selection.\n", "\n", "Because the scorer is just an estimator attribute, you can change it between\n", "calls to `partial_fit`. This notebook:\n", "\n", "1. Fits a model on an imbalanced binary problem using AUROC.\n", "2. Locks every internal node of the best program, leaving only the leaves and\n", " the weights free to change.\n", "3. Switches the scorer to average precision and calls `partial_fit`.\n", "4. Compares all metrics before and after the switch." ] }, { "cell_type": "code", "execution_count": 1, "id": "37c250d0", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:04.638378Z", "iopub.status.busy": "2026-09-22T19:12:04.638227Z", "iopub.status.idle": "2026-09-22T19:12:05.855283Z", "shell.execute_reply": "2026-09-22T19:12:05.854853Z" } }, "outputs": [], "source": [ "import numpy as np\n", "import pandas as pd\n", "import graphviz\n", "\n", "from sklearn.datasets import make_classification\n", "from sklearn.model_selection import train_test_split\n", "from sklearn.metrics import (log_loss, accuracy_score, balanced_accuracy_score,\n", " precision_score, recall_score, roc_auc_score,\n", " average_precision_score)\n", "\n", "from pybrush import BrushClassifier" ] }, { "cell_type": "markdown", "id": "db851330", "metadata": {}, "source": [ "## 1. An imbalanced binary problem\n", "\n", "AUROC and average precision disagree the most when positives are rare: AUROC\n", "rewards ranking negatives correctly, while average precision focuses on how\n", "clean the top of the ranking is." ] }, { "cell_type": "code", "execution_count": 2, "id": "c6262431", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:05.856401Z", "iopub.status.busy": "2026-09-22T19:12:05.856312Z", "iopub.status.idle": "2026-09-22T19:12:05.865082Z", "shell.execute_reply": "2026-09-22T19:12:05.864748Z" } }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "train prevalence: 0.109\n", "test prevalence: 0.107\n" ] } ], "source": [ "X, y = make_classification(\n", " n_samples=1000, n_features=6, n_informative=4, n_redundant=1,\n", " weights=[0.9, 0.1], class_sep=0.8, flip_y=0.02, random_state=42,\n", ")\n", "\n", "X_train, X_test, y_train, y_test = train_test_split(\n", " X, y, test_size=0.3, stratify=y, random_state=42)\n", "\n", "print('train prevalence:', y_train.mean().round(3))\n", "print('test prevalence: ', y_test.mean().round(3))" ] }, { "cell_type": "code", "execution_count": 3, "id": "ed5afaf0", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:05.865913Z", "iopub.status.busy": "2026-09-22T19:12:05.865867Z", "iopub.status.idle": "2026-09-22T19:12:05.867551Z", "shell.execute_reply": "2026-09-22T19:12:05.867314Z" } }, "outputs": [], "source": [ "def all_metrics(est, X, y):\n", " proba = est.predict_proba(X)[:, 1]\n", " pred = est.predict(X)\n", " return {\n", " 'log_loss': log_loss(y, proba),\n", " 'accuracy': accuracy_score(y, pred),\n", " 'balanced_accuracy': balanced_accuracy_score(y, pred),\n", " 'precision': precision_score(y, pred, zero_division=0),\n", " 'recall': recall_score(y, pred, zero_division=0),\n", " 'roc_auc': roc_auc_score(y, proba),\n", " 'average_precision': average_precision_score(y, proba),\n", " }\n", "\n", "def report(est, label):\n", " return pd.DataFrame({\n", " (label, 'train'): all_metrics(est, X_train, y_train),\n", " (label, 'test'): all_metrics(est, X_test, y_test),\n", " })" ] }, { "cell_type": "markdown", "id": "c072f652", "metadata": {}, "source": [ "## 2. Fit using AUROC as the scorer" ] }, { "cell_type": "code", "execution_count": 4, "id": "bfa036f8", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:05.868283Z", "iopub.status.busy": "2026-09-22T19:12:05.868241Z", "iopub.status.idle": "2026-09-22T19:12:14.024733Z", "shell.execute_reply": "2026-09-22T19:12:14.024331Z" } }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "scorer: roc_auc\n", "internal scorer value (train): 0.9027\n", "model: Logistic(Add(0.98,Add(-1.05*Mul(-1.32*x_4,-1.01*x_2),x_0)))\n" ] }, { "data": { "text/html": [ "
\n", "\n", "\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "
roc_auc
traintest
log_loss0.3982100.387116
accuracy0.8100000.843333
balanced_accuracy0.8009870.871035
precision0.3389830.397260
recall0.7894740.906250
roc_auc0.8869560.938666
average_precision0.6782970.789859
\n", "
" ], "text/plain": [ " roc_auc \n", " train test\n", "log_loss 0.398210 0.387116\n", "accuracy 0.810000 0.843333\n", "balanced_accuracy 0.800987 0.871035\n", "precision 0.338983 0.397260\n", "recall 0.789474 0.906250\n", "roc_auc 0.886956 0.938666\n", "average_precision 0.678297 0.789859" ] }, "execution_count": 4, "metadata": {}, "output_type": "execute_result" } ], "source": [ "est = BrushClassifier(\n", " functions=['SplitBest', 'Add', 'Sub', 'Mul', 'Div', 'Logabs', 'Exp'],\n", " scorer='roc_auc',\n", " pop_size=200,\n", " max_gens=30,\n", " max_depth=5,\n", " max_size=20,\n", " random_state=42,\n", " verbosity=0,\n", ")\n", "est.fit(X_train, y_train)\n", "\n", "print('scorer:', est.parameters_.scorer)\n", "print('internal scorer value (train):', round(est.best_estimator_.fitness.loss, 4))\n", "print('model:', est.best_estimator_.get_model())\n", "\n", "before = report(est, 'roc_auc')\n", "before" ] }, { "cell_type": "markdown", "id": "3aa6132b", "metadata": {}, "source": [ "## 3. Lock the internal nodes, switch to average precision, and refit\n", "\n", "A `lock_nodes_depth` larger than the tree depth locks every operator and split\n", "in the program. With `keep_leaves_unlocked=True`, the leaves (features and\n", "constants) stay free, so the search can still swap a leaf or grow it into a\n", "small subtree. With `keep_current_weights=False`, the weights are re-optimized.\n", "\n", "We then set `est.scorer = 'average_precision_score'` before calling\n", "`partial_fit`. The new scorer is picked up by the engine, so every program is\n", "re-evaluated, selected, and archived by average precision from here on." ] }, { "cell_type": "code", "execution_count": 5, "id": "b33832db", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:14.025731Z", "iopub.status.busy": "2026-09-22T19:12:14.025682Z", "iopub.status.idle": "2026-09-22T19:12:20.981227Z", "shell.execute_reply": "2026-09-22T19:12:20.980941Z" } }, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ "scorer: average_precision_score\n", "internal scorer value (train): 0.917\n", "model before: Logistic(Add(0.98,Add(-1.05*Mul(-1.32*x_4,-1.01*x_2),x_0)))\n", "model after: Logistic(Add(-2.77,Add(-1.53*Mul(x_4,x_2),3.87*Exp(0.42*x_0))))\n" ] } ], "source": [ "structure_before = est.best_estimator_.get_model()\n", "dot_before = est.best_estimator_.get_model('dot')\n", "\n", "est.scorer = 'average_precision_score'\n", "est.partial_fit(\n", " X_train, y_train,\n", " lock_nodes_depth=est.max_depth + 1, # deeper than any tree: lock every internal node\n", " keep_leaves_unlocked=True, # ...but leave the leaves free to change\n", " keep_current_weights=False, # weights can still be optimized\n", ")\n", "\n", "print('scorer:', est.parameters_.scorer)\n", "print('internal scorer value (train):', round(est.best_estimator_.fitness.loss, 4))\n", "print('model before:', structure_before)\n", "print('model after: ', est.best_estimator_.get_model())\n", "\n", "after = report(est, 'average_precision_score')" ] }, { "cell_type": "markdown", "id": "ed1c72d3", "metadata": {}, "source": [ "The two programs side by side. The locked internal nodes are the same in\n", "both, and only the leaves (and the weights) differ." ] }, { "cell_type": "code", "execution_count": 6, "id": "f7a39b67", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:20.982237Z", "iopub.status.busy": "2026-09-22T19:12:20.982178Z", "iopub.status.idle": "2026-09-22T19:12:21.166282Z", "shell.execute_reply": "2026-09-22T19:12:21.165925Z" } }, "outputs": [ { "data": { "text/html": [ "\n", "
\n", "

Before (AUROC)

\n", "\n", "\n", "\n", "\n", "\n", "G\n", "\n", "^ split feature fixed, * split threshold fixed\n", "\n", "\n", "1473c7000\n", "\n", "Logistic\n", "\n", "\n", "\n", "105024bc0\n", "\n", "Add\n", "\n", "\n", "\n", "1473c7000->105024bc0\n", "\n", "\n", "\n", "\n", "\n", "1473c70b0\n", "\n", "Add\n", "\n", "\n", "\n", "105024bc0->1473c70b0\n", "\n", "\n", "\n", "\n", "\n", "105024bc0Offset\n", "\n", "0.98\n", "\n", "\n", "\n", "105024bc0->105024bc0Offset\n", "\n", "\n", "\n", "\n", "\n", "1473884b0\n", "\n", "Mul\n", "\n", "\n", "\n", "1473c70b0->1473884b0\n", "\n", "\n", "-1.05\n", "\n", "\n", "\n", "x_0\n", "\n", "x_0\n", "\n", "\n", "\n", "1473c70b0->x_0\n", "\n", "\n", "\n", "\n", "\n", "x_4\n", "\n", "x_4\n", "\n", "\n", "\n", "1473884b0->x_4\n", "\n", "\n", "-1.32\n", "\n", "\n", "\n", "x_2\n", "\n", "x_2\n", "\n", "\n", "\n", "1473884b0->x_2\n", "\n", "\n", "-1.01\n", "\n", "\n", "\n", "
\n", "

After (average precision)

\n", "\n", "\n", "\n", "\n", "\n", "G\n", "\n", "^ split feature fixed, * split threshold fixed\n", "\n", "\n", "15f54b570\n", "\n", "Logistic\n", "\n", "\n", "\n", "15f54ba00\n", "\n", "Add\n", "\n", "\n", "\n", "15f54b570->15f54ba00\n", "\n", "\n", "\n", "\n", "\n", "15f5c09c0\n", "\n", "Add\n", "\n", "\n", "\n", "15f54ba00->15f5c09c0\n", "\n", "\n", "\n", "\n", "\n", "15f54ba00Offset\n", "\n", "-2.77\n", "\n", "\n", "\n", "15f54ba00->15f54ba00Offset\n", "\n", "\n", "\n", "\n", "\n", "15f5c0a70\n", "\n", "Mul\n", "\n", "\n", "\n", "15f5c09c0->15f5c0a70\n", "\n", "\n", "-1.53\n", "\n", "\n", "\n", "15f5f0460\n", "\n", "Exp\n", "\n", "\n", "\n", "15f5c09c0->15f5f0460\n", "\n", "\n", "3.87\n", "\n", "\n", "\n", "x_4\n", "\n", "x_4\n", "\n", "\n", "\n", "15f5c0a70->x_4\n", "\n", "\n", "\n", "\n", "\n", "x_2\n", "\n", "x_2\n", "\n", "\n", "\n", "15f5c0a70->x_2\n", "\n", "\n", "\n", "\n", "\n", "x_0\n", "\n", "x_0\n", "\n", "\n", "\n", "15f5f0460->x_0\n", "\n", "\n", "0.42\n", "\n", "\n", "\n", "
\n", "
\n" ], "text/plain": [ "" ] }, "execution_count": 6, "metadata": {}, "output_type": "execute_result" } ], "source": [ "from IPython.display import HTML\n", "\n", "def svg(dot):\n", " return graphviz.Source(dot).pipe(format='svg').decode()\n", "\n", "HTML(f\"\"\"\n", "
\n", "

Before (AUROC)

{svg(dot_before)}
\n", "

After (average precision)

{svg(est.best_estimator_.get_model('dot'))}
\n", "
\n", "\"\"\")" ] }, { "cell_type": "markdown", "id": "26aa005f", "metadata": {}, "source": [ "## 4. Compare every metric before and after the switch" ] }, { "cell_type": "code", "execution_count": 7, "id": "32d527ab", "metadata": { "execution": { "iopub.execute_input": "2026-09-22T19:12:21.167352Z", "iopub.status.busy": "2026-09-22T19:12:21.167284Z", "iopub.status.idle": "2026-09-22T19:12:21.173975Z", "shell.execute_reply": "2026-09-22T19:12:21.173673Z" } }, "outputs": [ { "data": { "text/html": [ "
\n", "\n", "\n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", " \n", "
roc_aucaverage_precision_scoredelta
traintesttraintesttraintest
log_loss0.39820.38710.38190.3787-0.0163-0.0084
accuracy0.81000.84330.84000.86330.03000.0200
balanced_accuracy0.80100.87100.81780.88220.01680.0112
precision0.33900.39730.38460.43280.04560.0356
recall0.78950.90620.78950.90620.00000.0000
roc_auc0.88700.93870.88920.93970.00220.0010
average_precision0.67830.78990.70410.79830.02580.0085
\n", "
" ], "text/plain": [ " roc_auc average_precision_score delta \\\n", " train test train test train \n", "log_loss 0.3982 0.3871 0.3819 0.3787 -0.0163 \n", "accuracy 0.8100 0.8433 0.8400 0.8633 0.0300 \n", "balanced_accuracy 0.8010 0.8710 0.8178 0.8822 0.0168 \n", "precision 0.3390 0.3973 0.3846 0.4328 0.0456 \n", "recall 0.7895 0.9062 0.7895 0.9062 0.0000 \n", "roc_auc 0.8870 0.9387 0.8892 0.9397 0.0022 \n", "average_precision 0.6783 0.7899 0.7041 0.7983 0.0258 \n", "\n", " \n", " test \n", "log_loss -0.0084 \n", "accuracy 0.0200 \n", "balanced_accuracy 0.0112 \n", "precision 0.0356 \n", "recall 0.0000 \n", "roc_auc 0.0010 \n", "average_precision 0.0085 " ] }, "execution_count": 7, "metadata": {}, "output_type": "execute_result" } ], "source": [ "comparison = pd.concat([before, after], axis=1)\n", "comparison[('delta', 'train')] = comparison[('average_precision_score', 'train')] - comparison[('roc_auc', 'train')]\n", "comparison[('delta', 'test')] = comparison[('average_precision_score', 'test')] - comparison[('roc_auc', 'test')]\n", "comparison.round(4)" ] } ], "metadata": { "kernelspec": { "display_name": "brush", "language": "python", "name": "python3" }, "language_info": { "codemirror_mode": { "name": "ipython", "version": 3 }, "file_extension": ".py", "mimetype": "text/x-python", "name": "python", "nbconvert_exporter": "python", "pygments_lexer": "ipython3", "version": "3.13.14" } }, "nbformat": 4, "nbformat_minor": 5 }