diff --git a/docs/guide/index.md b/docs/guide/index.md index f85cc9550..eb09ae2aa 100644 --- a/docs/guide/index.md +++ b/docs/guide/index.md @@ -12,9 +12,10 @@ Brush mostly consists of these components: data search_space working_with_programs +multiclassification json saving_loading_populations locking_mechanism archive deap -``` \ No newline at end of file +``` diff --git a/docs/guide/locking_mechanism.ipynb b/docs/guide/locking_mechanism.ipynb index 5ecd47e70..1eeb27720 100644 --- a/docs/guide/locking_mechanism.ipynb +++ b/docs/guide/locking_mechanism.ipynb @@ -54,7 +54,7 @@ }, { "cell_type": "code", - "execution_count": 3, + "execution_count": 8, "id": "d4128f4b", "metadata": {}, "outputs": [ @@ -63,7 +63,9 @@ "output_type": "stream", "text": [ "base shape: (500, 3) (500,)\n", - "shifted shape: (500, 3) (500,)\n" + "shifted shape: (500, 3) (500,)\n", + "y base mean: -0.07736613840644213\n", + "y shifted mean: 3.7191083728546763\n" ] } ], @@ -88,7 +90,9 @@ "])\n", "\n", "print('base shape:', X_base.shape, y_base.shape)\n", - "print('shifted shape:', X_shifted.shape, y_shifted.shape)" + "print('shifted shape:', X_shifted.shape, y_shifted.shape)\n", + "print('y base mean: ', np.mean(y_base))\n", + "print('y shifted mean:', np.mean(y_shifted))" ] }, { @@ -112,23 +116,23 @@ "output_type": "stream", "text": [ "Completed 100% [====================]\n", - "stage 1 base mse: 2.024100915882261\n", - "stage 1 shifted mse: 6.257131345427466\n", + "stage 1 base mse: 0.9389598521387683\n", + "stage 1 shifted mse: 12.693906935930773\n", "stage 1 model:\n", - "Sub(Mul(Add(x_1,If(x_1>=-2.61,x_0,-4.43)),Sub(1.47*x_2,1.31)),-3.38*x_0)\n" + "If(x_0>=0.24,If(x_2>=-0.04,3.90,-3.03*x_1),If(x_2>=0.04,4.09*x_0,-2.98*x_1))\n" ] } ], "source": [ "est = BrushRegressor(\n", " functions=['SplitOn', 'SplitBest', 'Mul', 'Add', 'Sub'],\n", - " pop_size=100,\n", - " max_gens=25,\n", - " max_depth=10,\n", - " max_size=24,\n", + " pop_size=500,\n", + " max_gens=20,\n", + " max_depth=5,\n", + " max_size=10,\n", " start_from_decision_trees=True,\n", " constants_simplification=True,\n", - " inexact_simplification=False,\n", + " inexact_simplification=True,\n", " verbosity=1,\n", " random_state=7,\n", ")\n", @@ -163,19 +167,15 @@ "text": [ "locked the top two levels of the current best estimator\n", "locked model:\n", - "Sub(Mul(Add(x_1,If(x_1>=-2.61,x_0,-4.43)),Sub(1.47*x_2,1.31)),-3.38*x_0)\n", + "If(x_0>=0.24,If(x_2>=-0.04,3.90,-3.03*x_1),If(x_2>=0.04,4.09*x_0,-2.98*x_1))\n", "locked tree:\n", - "Sub\n", - "|- Mul\n", - "| |- Add\n", - "| | |- x_1\n", - "| | |- If(x_1>=-2.61)\n", - "| | | |- x_0\n", - "| | | |- -4.43\n", - "| |- Sub\n", - "| | |- 1.47*x_2\n", - "| | |- 1.31\n", - "|- -3.38*x_0\n" + "If(x_0>=0.24)\n", + "|- If(x_2>=-0.04)\n", + "| |- 3.90\n", + "| |- -3.03*x_1\n", + "|- If(x_2>=0.04)\n", + "| |- 4.09*x_0\n", + "| |- -2.98*x_1\n" ] }, { @@ -187,141 +187,111 @@ "\n", "\n", - "\n", - "\n", + "\n", + "\n", "G\n", - "\n", - "^ split feature fixed, * split threshold fixed\n", - "\n", + "\n", + "^ split feature fixed, * split threshold fixed\n", + "\n", "\n", - "14f96e850\n", - "\n", - "Sub\n", + "y\n", + "\n", + "y\n", "\n", - "\n", + "\n", "\n", - "14f96e7a0\n", - "\n", - "Mul\n", + "1275d86b0\n", + "\n", + "x_0 >= 0.24*?\n", "\n", - "\n", + "\n", "\n", - "14f96e850->14f96e7a0\n", - "\n", - "\n", + "y->1275d86b0\n", + "\n", + "\n", + "0.24\n", "\n", - "\n", + "\n", "\n", - "x_0\n", - "\n", - "x_0\n", + "1275d6d30\n", + "\n", + "x_2 >= -0.04*?\n", "\n", - "\n", + "\n", "\n", - "14f96e850->x_0\n", - "\n", - "\n", - "-3.38\n", + "1275d86b0->1275d6d30\n", + "\n", + "\n", + "Y\n", "\n", - "\n", + "\n", "\n", - "14f971350\n", - "\n", - "Add\n", + "17891b060\n", + "\n", + "x_2 >= 0.04*?\n", "\n", - "\n", + "\n", "\n", - "14f96e7a0->14f971350\n", - "\n", - "\n", + "1275d86b0->17891b060\n", + "\n", + "\n", + "N\n", "\n", - "\n", + "\n", "\n", - "14f971820\n", - "\n", - "Sub\n", + "1275c4ff0\n", + "\n", + "3.90\n", "\n", - "\n", + "\n", "\n", - "14f96e7a0->14f971820\n", - "\n", - "\n", + "1275d6d30->1275c4ff0\n", + "\n", + "\n", + "Y\n", "\n", "\n", "\n", "x_1\n", - "\n", - "x_1\n", + "\n", + "x_1\n", "\n", - "\n", + "\n", "\n", - "14f971350->x_1\n", - "\n", - "\n", - "\n", - "\n", - "\n", - "14f9714b0\n", - "\n", - "x_1 >= -2.61?\n", - "\n", - "\n", - "\n", - "14f971350->14f9714b0\n", - "\n", - "\n", - "\n", - "\n", - "\n", - "x_2\n", - "\n", - "x_2\n", + "1275d6d30->x_1\n", + "\n", + "\n", + "-3.03\n", + "N\n", "\n", - "\n", - "\n", - "14f971820->x_2\n", - "\n", - "\n", - "1.47\n", - "\n", - "\n", - "\n", - "14f971980\n", - "\n", - "1.31\n", - "\n", - "\n", - "\n", - "14f971820->14f971980\n", - "\n", - "\n", - "\n", - "\n", + "\n", "\n", - "14f9714b0->x_0\n", - "\n", - "\n", - "Y\n", + "17891b060->x_1\n", + "\n", + "\n", + "-2.98\n", + "N\n", "\n", - "\n", - "\n", - "14f971610\n", - "\n", - "-4.43\n", + "\n", + "\n", + "x_0\n", + "\n", + "x_0\n", "\n", - "\n", - "\n", - "14f9714b0->14f971610\n", - "\n", - "\n", - "N\n", + "\n", + "\n", + "17891b060->x_0\n", + "\n", + "\n", + "4.09\n", + "Y\n", "\n", "\n", "\n" ], "text/plain": [ - "" + "" ] }, "execution_count": 5, @@ -331,7 +301,7 @@ ], "source": [ "est.best_estimator_.program.lock_nodes(\n", - " 3,\n", + " 2,\n", " keep_leaves_unlocked=True,\n", " keep_current_weights=True,\n", ")\n", @@ -365,24 +335,18 @@ "output_type": "stream", "text": [ "Completed 100% [====================]\n", - "stage 2 base mse: 3.6734543201338834\n", - "stage 2 shifted mse: 3.8054243085275217\n", + "stage 2 base mse: 2.4325052852682343\n", + "stage 2 shifted mse: 0.9083711174803811\n", "stage 2 model:\n", - "Sub(Mul(Add(0.85*Sub(1.81*x_0,-2.40*x_1),0.01),Sub(x_2,-1.40)),0.88*Sub(-0.22,-6.02*x_1))\n", + "If(x_0>=0.24,If(x_2>=-0.04,4.06*x_0,-5.26*x_1),If(x_2>=0.04,3.96*x_0,-4.84*x_1))\n", "stage 2 tree:\n", - "Sub\n", - "|- Mul\n", - "| |- Add\n", - "| | |- 0.85*Sub\n", - "| | | |- 1.81*x_0\n", - "| | | |- -2.40*x_1\n", - "| | |- 0.01\n", - "| |- Sub\n", - "| | |- x_2\n", - "| | |- -1.40\n", - "|- 0.88*Sub\n", - "| |- -0.22\n", - "| |- -6.02*x_1\n" + "If(x_0>=0.24)\n", + "|- If(x_2>=-0.04)\n", + "| |- 4.06*x_0\n", + "| |- -5.26*x_1\n", + "|- If(x_2>=0.04)\n", + "| |- 3.96*x_0\n", + "| |- -4.84*x_1\n" ] }, { @@ -394,166 +358,106 @@ "\n", "\n", - "\n", - "\n", + "\n", + "\n", "G\n", - "\n", - "^ split feature fixed, * split threshold fixed\n", - "\n", + "\n", + "^ split feature fixed, * split threshold fixed\n", + "\n", "\n", - "14f984ee0\n", - "\n", - "Sub\n", + "y\n", + "\n", + "y\n", "\n", - "\n", + "\n", "\n", - "14f96c1a0\n", - "\n", - "Mul\n", + "105973df0\n", + "\n", + "x_0 >= 0.24*?\n", "\n", - "\n", + "\n", "\n", - "14f984ee0->14f96c1a0\n", - "\n", - "\n", + "y->105973df0\n", + "\n", + "\n", + "0.24\n", "\n", - "\n", + "\n", "\n", - "14f969cf0\n", - "\n", - "Sub\n", + "105997a90\n", + "\n", + "x_2 >= -0.04*?\n", "\n", - "\n", + "\n", "\n", - "14f984ee0->14f969cf0\n", - "\n", - "\n", - "0.88\n", + "105973df0->105997a90\n", + "\n", + "\n", + "Y\n", "\n", - "\n", + "\n", "\n", - "14f96f6c0\n", - "\n", - "Add\n", + "1059e3ca0\n", + "\n", + "x_2 >= 0.04*?\n", "\n", - "\n", + "\n", "\n", - "14f96c1a0->14f96f6c0\n", - "\n", - "\n", + "105973df0->1059e3ca0\n", + "\n", + "\n", + "N\n", "\n", - "\n", + "\n", "\n", - "14f983360\n", - "\n", - "Sub\n", + "x_0\n", + "\n", + "x_0\n", "\n", - "\n", + "\n", "\n", - "14f96c1a0->14f983360\n", - "\n", - "\n", + "105997a90->x_0\n", + "\n", + "\n", + "4.06\n", + "Y\n", "\n", "\n", - "\n", - "x_1\n", - "\n", - "x_1\n", - "\n", - "\n", - "\n", - "14f969cf0->x_1\n", - "\n", - "\n", - "-6.02\n", - "\n", - "\n", - "\n", - "14f98af40\n", - "\n", - "-0.22\n", - "\n", - "\n", - "\n", - "14f969cf0->14f98af40\n", - "\n", - "\n", - "\n", - "\n", "\n", - "14f96df30\n", - "\n", - "Sub\n", + "x_1\n", + "\n", + "x_1\n", "\n", - "\n", + "\n", "\n", - "14f96f6c0->14f96df30\n", - "\n", - "\n", - "0.85\n", + "105997a90->x_1\n", + "\n", + "\n", + "-5.26\n", + "N\n", "\n", - "\n", - "\n", - "14f987090\n", - "\n", - "0.01\n", - "\n", - "\n", + "\n", "\n", - "14f96f6c0->14f987090\n", - "\n", - "\n", - "\n", - "\n", - "\n", - "x_2\n", - "\n", - "x_2\n", - "\n", - "\n", - "\n", - "14f983360->x_2\n", - "\n", - "\n", + "1059e3ca0->x_0\n", + "\n", + "\n", + "3.96\n", + "Y\n", "\n", - "\n", - "\n", - "14f969140\n", - "\n", - "-1.40\n", - "\n", - "\n", - "\n", - "14f983360->14f969140\n", - "\n", - "\n", - "\n", - "\n", - "\n", - "x_0\n", - "\n", - "x_0\n", - "\n", - "\n", + "\n", "\n", - "14f96df30->x_0\n", - "\n", - "\n", - "1.81\n", - "\n", - "\n", - "\n", - "14f96df30->x_1\n", - "\n", - "\n", - "-2.40\n", + "1059e3ca0->x_1\n", + "\n", + "\n", + "-4.84\n", + "N\n", "\n", "\n", "\n" ], "text/plain": [ - "" + "" ] }, "execution_count": 6, @@ -565,7 +469,7 @@ "est.partial_fit(\n", " X_shifted,\n", " y_shifted,\n", - " lock_nodes_depth=3,\n", + " lock_nodes_depth=2,\n", " keep_leaves_unlocked=True,\n", " keep_current_weights=True,\n", ")\n", @@ -600,24 +504,24 @@ "output_type": "stream", "text": [ "Completed 100% [====================]\n", - "stage 3 base mse: 3.648386877972349\n", - "stage 3 shifted mse: 3.7962339333011497\n", + "stage 3 base mse: 5.744836281693709\n", + "stage 3 shifted mse: 3.5116005509724473\n", "stage 3 unlocked model:\n", - "Sub(Mul(Add(Sub(1.57*x_0,-2.03*x_1),-0.05*x_2),Sub(x_2,-1.38)),0.93*Sub(-0.29,-5.58*x_1))\n", + "If(x_0>=2.05,If(x_2>=-0.12,3.98*x_0,0.40*Sub(x_0,12.50*x_1)),If(x_1>=-0.59,2.06*Sub(x_0,x_1),If(x_2>=0.00,4.17*x_0,-5.12*x_1)))\n", "stage 3 unlocked tree:\n", - "Sub\n", - "|- Mul\n", - "| |- Add\n", - "| | |- Sub\n", - "| | | |- 1.57*x_0\n", - "| | | |- -2.03*x_1\n", - "| | |- -0.05*x_2\n", - "| |- Sub\n", - "| | |- x_2\n", - "| | |- -1.38\n", - "|- 0.93*Sub\n", - "| |- -0.29\n", - "| |- -5.58*x_1\n" + "If(x_0>=2.05)\n", + "|- If(x_2>=-0.12)\n", + "| |- 3.98*x_0\n", + "| |- 0.40*Sub\n", + "| | |- x_0\n", + "| | |- 12.50*x_1\n", + "|- If(x_1>=-0.59)\n", + "| |- 2.06*Sub\n", + "| | |- x_0\n", + "| | |- x_1\n", + "| |- If(x_2>=0.00)\n", + "| | |- 4.17*x_0\n", + "| | |- -5.12*x_1\n" ] }, { @@ -629,160 +533,164 @@ "\n", "\n", - "\n", - "\n", + "\n", + "\n", "G\n", - "\n", - "^ split feature fixed, * split threshold fixed\n", - "\n", + "\n", + "^ split feature fixed, * split threshold fixed\n", + "\n", "\n", - "10ec2d9d0\n", - "\n", - "Sub\n", + "y\n", + "\n", + "y\n", "\n", - "\n", + "\n", "\n", - "10d8865a0\n", - "\n", - "Mul\n", + "102e2ed40\n", + "\n", + "x_0 >= 2.05?\n", "\n", - "\n", + "\n", "\n", - "10ec2d9d0->10d8865a0\n", - "\n", - "\n", + "y->102e2ed40\n", + "\n", + "\n", + "2.05\n", "\n", - "\n", + "\n", "\n", - "10d8806a0\n", - "\n", - "Sub\n", + "102e5edc0\n", + "\n", + "x_2 >= -0.12?\n", "\n", - "\n", + "\n", "\n", - "10ec2d9d0->10d8806a0\n", - "\n", - "\n", - "0.93\n", + "102e2ed40->102e5edc0\n", + "\n", + "\n", + "Y\n", "\n", - "\n", + "\n", "\n", - "14f971580\n", - "\n", - "Add\n", + "127489710\n", + "\n", + "x_1 >= -0.59?\n", "\n", - "\n", + "\n", "\n", - "10d8865a0->14f971580\n", - "\n", - "\n", + "102e2ed40->127489710\n", + "\n", + "\n", + "N\n", "\n", - "\n", + "\n", "\n", - "10d832110\n", - "\n", - "Sub\n", + "x_0\n", + "\n", + "x_0\n", "\n", - "\n", + "\n", "\n", - "10d8865a0->10d832110\n", - "\n", - "\n", - "\n", - "\n", - "\n", - "x_1\n", - "\n", - "x_1\n", + "102e5edc0->x_0\n", + "\n", + "\n", + "3.98\n", + "Y\n", "\n", - "\n", - "\n", - "10d8806a0->x_1\n", - "\n", - "\n", - "-5.58\n", - "\n", - "\n", - "\n", - "10d8a8f30\n", - "\n", - "-0.29\n", - "\n", - "\n", - "\n", - "10d8806a0->10d8a8f30\n", - "\n", - "\n", - "\n", - "\n", + "\n", "\n", - "10ec1b4f0\n", - "\n", - "Sub\n", + "1274e7de0\n", + "\n", + "Sub\n", "\n", - "\n", + "\n", "\n", - "14f971580->10ec1b4f0\n", - "\n", - "\n", + "102e5edc0->1274e7de0\n", + "\n", + "\n", + "0.40\n", + "N\n", "\n", - "\n", - "\n", - "x_2\n", - "\n", - "x_2\n", + "\n", + "\n", + "102e61a80\n", + "\n", + "Sub\n", "\n", - "\n", - "\n", - "14f971580->x_2\n", - "\n", - "\n", - "-0.05\n", + "\n", + "\n", + "127489710->102e61a80\n", + "\n", + "\n", + "2.06\n", + "Y\n", + "\n", + "\n", + "\n", + "10595fde0\n", + "\n", + "x_2 >= 0.00?\n", "\n", - "\n", + "\n", "\n", - "10d832110->x_2\n", - "\n", - "\n", + "127489710->10595fde0\n", + "\n", + "\n", + "N\n", "\n", - "\n", - "\n", - "10b5774f0\n", - "\n", - "-1.38\n", + "\n", + "\n", + "1274e7de0->x_0\n", + "\n", + "\n", "\n", - "\n", + "\n", + "\n", + "x_1\n", + "\n", + "x_1\n", + "\n", + "\n", + "\n", + "1274e7de0->x_1\n", + "\n", + "\n", + "12.50\n", + "\n", + "\n", "\n", - "10d832110->10b5774f0\n", - "\n", - "\n", + "102e61a80->x_0\n", + "\n", + "\n", "\n", - "\n", - "\n", - "x_0\n", - "\n", - "x_0\n", + "\n", + "\n", + "102e61a80->x_1\n", + "\n", + "\n", "\n", - "\n", - "\n", - "10ec1b4f0->x_0\n", - "\n", - "\n", - "1.57\n", + "\n", + "\n", + "10595fde0->x_0\n", + "\n", + "\n", + "4.17\n", + "Y\n", "\n", - "\n", - "\n", - "10ec1b4f0->x_1\n", - "\n", - "\n", - "-2.03\n", + "\n", + "\n", + "10595fde0->x_1\n", + "\n", + "\n", + "-5.12\n", + "N\n", "\n", "\n", "\n" ], "text/plain": [ - "" + "" ] }, "execution_count": 7, diff --git a/docs/guide/multiclassification.ipynb b/docs/guide/multiclassification.ipynb new file mode 100644 index 000000000..185dc0b9e --- /dev/null +++ b/docs/guide/multiclassification.ipynb @@ -0,0 +1,522 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "# Multiclass classification\n", + "\n", + "This example fits a three-class model, prints its tree representation, and renders the same tree with Graphviz." + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "4859319e", + "metadata": {}, + "outputs": [], + "source": [ + "from sklearn.datasets import make_classification\n", + "from sklearn.model_selection import train_test_split\n", + "from sklearn.metrics import accuracy_score\n", + "from pybrush import BrushClassifier\n", + "import graphviz\n", + "\n", + "X, y = make_classification(\n", + " n_samples=180, n_features=6, n_informative=5, n_redundant=0,\n", + " n_classes=3, n_clusters_per_class=1, random_state=42,\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", + ")" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "id": "b3d1aef8", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Completed 100% [====================]\n", + "Test accuracy: 0.7962962962962963\n", + "Probability rows sum to: [1. 1. 1.]\n" + ] + } + ], + "source": [ + "model = BrushClassifier(\n", + " pop_size=100, max_gens=100, max_size=40, max_depth=10,\n", + " num_islands=1, random_state=42, verbosity=1,\n", + ")\n", + "model.fit(X_train, y_train)\n", + "\n", + "print('Test accuracy:', accuracy_score(y_test, model.predict(X_test)))\n", + "print('Probability rows sum to:', model.predict_proba(X_test)[:3].sum(axis=1))" + ] + }, + { + "cell_type": "code", + "execution_count": 3, + "id": "56a9995f", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Softmax\n", + "|- x_0\n", + "|- Mean\n", + "| |- 12.73*x_2\n", + "| |- Min\n", + "| | |- 55.80*x_5\n", + "| | |- -7.67\n", + "| | |- -2.40\n", + "| | |- 1.00\n", + "| |- -14.89*x_4\n", + "| |- 20.88*x_3\n", + "|- Sum\n", + "| |- -1.35*x_3\n", + "| |- -0.76*x_2\n", + "| |- 0.51*x_0\n" + ] + } + ], + "source": [ + "print(model.best_estimator_.program.get_model('tree'))" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "id": "a46f2e0d", + "metadata": {}, + "outputs": [ + { + "data": { + "image/svg+xml": [ + "\n", + "\n", + "\n", + "\n", + "\n", + "\n", + "G\n", + "\n", + "^ split feature fixed, * split threshold fixed\n", + "\n", + "\n", + "1641b0670\n", + "\n", + "Softmax\n", + "\n", + "\n", + "\n", + "x_0\n", + "\n", + "x_0\n", + "\n", + "\n", + "\n", + "1641b0670->x_0\n", + "\n", + "\n", + "class 0\n", + "\n", + "\n", + "\n", + "1641353b0\n", + "\n", + "Mean\n", + "\n", + "\n", + "\n", + "1641b0670->1641353b0\n", + "\n", + "\n", + "class 1\n", + "\n", + "\n", + "\n", + "164121200\n", + "\n", + "Sum\n", + "\n", + "\n", + "\n", + "1641b0670->164121200\n", + "\n", + "\n", + "class 2\n", + "\n", + "\n", + "\n", + "x_2\n", + "\n", + "x_2\n", + "\n", + "\n", + "\n", + "1641353b0->x_2\n", + "\n", + "\n", + "12.73\n", + "\n", + "\n", + "\n", + "164454d70\n", + "\n", + "Min\n", + "\n", + "\n", + "\n", + "1641353b0->164454d70\n", + "\n", + "\n", + "\n", + "\n", + "\n", + "x_4\n", + "\n", + "x_4\n", + "\n", + "\n", + "\n", + "1641353b0->x_4\n", + "\n", + "\n", + "-14.89\n", + "\n", + "\n", + "\n", + "x_3\n", + "\n", + "x_3\n", + "\n", + "\n", + "\n", + "1641353b0->x_3\n", + "\n", + "\n", + "20.88\n", + "\n", + "\n", + "\n", + "164121200->x_0\n", + "\n", + "\n", + "0.51\n", + "\n", + "\n", + "\n", + "164121200->x_2\n", + "\n", + "\n", + "-0.76\n", + "\n", + "\n", + "\n", + "164121200->x_3\n", + "\n", + "\n", + "-1.35\n", + "\n", + "\n", + "\n", + "x_5\n", + "\n", + "x_5\n", + "\n", + "\n", + "\n", + "164454d70->x_5\n", + "\n", + "\n", + "55.80\n", + "\n", + "\n", + "\n", + "164154f90\n", + "\n", + "-7.67\n", + "\n", + "\n", + "\n", + "164454d70->164154f90\n", + "\n", + "\n", + "\n", + "\n", + "\n", + "1641c5680\n", + "\n", + "-2.40\n", + "\n", + "\n", + "\n", + "164454d70->1641c5680\n", + "\n", + "\n", + "\n", + "\n", + "\n", + "16411a550\n", + "\n", + "1.00\n", + "\n", + "\n", + "\n", + "164454d70->16411a550\n", + "\n", + "\n", + "\n", + "\n", + "\n" + ], + "text/plain": [ + "" + ] + }, + "execution_count": 4, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "graphviz.Source(model.best_estimator_.program.get_model('dot'))" + ] + }, + { + "cell_type": "markdown", + "id": "a5c111df", + "metadata": {}, + "source": [ + "## Multiclass decision trees only\n", + "\n", + "Set `start_from_decision_trees=True` to restrict the initial model population to split-based class-logit branches. The Softmax root still combines one branch per class." + ] + }, + { + "cell_type": "code", + "execution_count": 14, + "id": "99bd6400", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Completed 15% [=== ]Split-only test accuracy: 0.7777777777777778\n", + "Softmax\n", + "|- 0.30*x_0\n", + "|- If(x_4>=-0.47)\n", + "| |- -1.53*x_4\n", + "| |- 3.39\n", + "|- If(x_4>=-0.47)\n", + "| |- 0.14*Add\n", + "| | |- -8.83*x_3\n", + "| | |- -5.89*x_2\n", + "| |- -0.43*x_2\n" + ] + } + ], + "source": [ + "split_model = BrushClassifier(\n", + " pop_size=100, max_gens=100, max_size=40, max_depth=5, max_stall=10,\n", + " start_from_decision_trees=True,\n", + " functions=['SplitOn', 'SplitBest', 'Add', 'Mul'],\n", + " num_islands=1, random_state=7, verbosity=1,\n", + ")\n", + "split_model.fit(X_train, y_train)\n", + "\n", + "print('Split-only test accuracy:', accuracy_score(y_test, split_model.predict(X_test)))\n", + "print(split_model.best_estimator_.program.get_model('tree'))" + ] + }, + { + "cell_type": "code", + "execution_count": 15, + "id": "3f700939", + "metadata": {}, + "outputs": [ + { + "data": { + "image/svg+xml": [ + "\n", + "\n", + "\n", + "\n", + "\n", + "\n", + "G\n", + "\n", + "^ split feature fixed, * split threshold fixed\n", + "\n", + "\n", + "1638d9460\n", + "\n", + "Softmax\n", + "\n", + "\n", + "\n", + "x_0\n", + "\n", + "x_0\n", + "\n", + "\n", + "\n", + "1638d9460->x_0\n", + "\n", + "\n", + "class 0\n", + "\n", + "\n", + "\n", + "1638cf8c0\n", + "\n", + "x_4 >= -0.47?\n", + "\n", + "\n", + "\n", + "1638d9460->1638cf8c0\n", + "\n", + "\n", + "class 1\n", + "\n", + "\n", + "\n", + "1638ccc20\n", + "\n", + "x_4 >= -0.47?\n", + "\n", + "\n", + "\n", + "1638d9460->1638ccc20\n", + "\n", + "\n", + "class 2\n", + "\n", + "\n", + "\n", + "x_4\n", + "\n", + "x_4\n", + "\n", + "\n", + "\n", + "1638cf8c0->x_4\n", + "\n", + "\n", + "-1.53\n", + "Y\n", + "\n", + "\n", + "\n", + "16384d160\n", + "\n", + "3.39\n", + "\n", + "\n", + "\n", + "1638cf8c0->16384d160\n", + "\n", + "\n", + "N\n", + "\n", + "\n", + "\n", + "1638ca800\n", + "\n", + "Add\n", + "\n", + "\n", + "\n", + "1638ccc20->1638ca800\n", + "\n", + "\n", + "0.14\n", + "Y\n", + "\n", + "\n", + "\n", + "x_2\n", + "\n", + "x_2\n", + "\n", + "\n", + "\n", + "1638ccc20->x_2\n", + "\n", + "\n", + "-0.43\n", + "N\n", + "\n", + "\n", + "\n", + "1638ca800->x_2\n", + "\n", + "\n", + "-5.89\n", + "\n", + "\n", + "\n", + "x_3\n", + "\n", + "x_3\n", + "\n", + "\n", + "\n", + "1638ca800->x_3\n", + "\n", + "\n", + "-8.83\n", + "\n", + "\n", + "\n" + ], + "text/plain": [ + "" + ] + }, + "execution_count": 15, + "metadata": {}, + "output_type": "execute_result" + } + ], + "source": [ + "graphviz.Source(split_model.best_estimator_.program.get_model('dot'))" + ] + } + ], + "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 +} diff --git a/pybrush/BrushEstimator.py b/pybrush/BrushEstimator.py index 5a74ac7b3..6ed3321a9 100644 --- a/pybrush/BrushEstimator.py +++ b/pybrush/BrushEstimator.py @@ -93,6 +93,12 @@ def fit(self, X, y): # Beyong this point, X is not a dataframe anymore X, y = check_X_y(X, y) + if self.mode == 'classification': + # The C++ core indexes probability columns with class labels. Keep + # that representation internal while preserving sklearn's original + # labels at the public API boundary. + self.classes_, y = np.unique(y, return_inverse=True) + y = y.astype(np.float32) self.data_ = self._make_data(X, y, feature_names=self.feature_names_, @@ -169,6 +175,14 @@ def partial_fit(self, X, y, *, assert self.feature_names_ == X.columns.to_list(), \ "Feature names must be the same as in data from previous fit" + if self.mode == 'classification': + labels = np.asarray(y) + indices = np.searchsorted(self.classes_, labels) + if np.any(indices >= len(self.classes_)) \ + or np.any(self.classes_[indices] != labels): + raise ValueError("partial_fit received a class not seen in fit") + y = indices.astype(np.float32) + new_data = self._make_data(X, y, feature_names=self.feature_names_, feature_types=self.feature_types_, @@ -228,7 +242,10 @@ def predict(self, X): validation_size=0.0, ) - return self.best_estimator_.program.predict(data) + prediction = np.asarray(self.best_estimator_.program.predict(data)) + if self.mode == 'classification': + return self.classes_[prediction.astype(int)] + return prediction def get_params(self, deep=True): out = dict() @@ -263,7 +280,8 @@ def _update_final_model(self, data=None): elif self.final_model_selection == "best_validation_ci": loss_f_dict = { # using sklearn metric, equivalent to what is used internally in brush "mse": mean_squared_error, - "log": log_loss, + "log": log_loss, + "multi_log": log_loss, "accuracy": accuracy_score, "balanced_accuracy": balanced_accuracy_score, "average_precision_score": average_precision_score @@ -274,11 +292,15 @@ def eval(ind, sample=None): if sample is None: sample = np.arange(len(y)) - if self.parameters_.scorer in ["log", "average_precision_score"]: + if self.parameters_.scorer in ["log", "multi_log", "average_precision_score"]: y_pred = np.array(ind.predict_proba(data)) else: # accuracy, balanced accuracy, or regression metrics y_pred = np.array(ind.predict(data)) + metric_kwargs = {} + if self.parameters_.scorer == "multi_log": + metric_kwargs["labels"] = np.arange(self.parameters_.n_classes) + # y_pred = np.nan_to_num(y_pred) # Protecting the evaluation # if user_defined, sample_weight is given by his custom weights. if @@ -300,9 +322,10 @@ def eval(ind, sample=None): # sample_weight will be indexed in the function call, so we use raw y. sample_weight = [support_weights[int(label)] for label in y] sample_weight = np.array(sample_weight) - return loss_f(y[sample], y_pred[sample], sample_weight=sample_weight[sample]) + return loss_f(y[sample], y_pred[sample], + sample_weight=sample_weight[sample], **metric_kwargs) else: # unbalanced metrics, ignoring weights - return loss_f(y[sample], y_pred[sample]) + return loss_f(y[sample], y_pred[sample], **metric_kwargs) np.random.seed(0) val_samples = [] @@ -427,4 +450,4 @@ class BrushRegressor(BrushEstimator, RegressorMixin): def __init__(self, **kwargs): kwargs.pop('mode', None) - super().__init__(mode='regression', **kwargs) \ No newline at end of file + super().__init__(mode='regression', **kwargs) diff --git a/pybrush/EstimatorInterface.py b/pybrush/EstimatorInterface.py index c3a9db103..afbca056c 100644 --- a/pybrush/EstimatorInterface.py +++ b/pybrush/EstimatorInterface.py @@ -77,7 +77,7 @@ class EstimatorInterface(): scorer : str, default None The metric to use for the "scorer" objective. If None, it will be set to "mse" for regression and "log" for binary classification. - Available options are `["mse", "log", "accuracy", "balanced_accuracy", "average_precision_score"]` + Available options are `["mse", "log", "multi_log", "accuracy", "balanced_accuracy", "average_precision_score"]` algorithm : {"nsga2island", "nsga2", "gaisland", "ga"}, default "nsga2" Which Evolutionary Algorithm framework to use to evolve the population. This is used only in DeapEstimators. @@ -350,10 +350,13 @@ def _wrap_parameters(self, y, **extra_kwargs): if self.mode == "regression": assert self.scorer in ['mse'], \ "Invalid scorer for the regression mode" + elif params.n_classes == 2: + assert self.scorer in ['log', 'balanced_accuracy', 'accuracy', + 'average_precision_score'], \ + "Invalid scorer for binary classification" else: - assert self.scorer in ['log', 'multi_log', 'balanced_accuracy', - 'accuracy', 'average_precision_score'], \ - "Invalid scorer for the classification mode" + assert self.scorer in ['multi_log', 'balanced_accuracy', 'accuracy'], \ + "Invalid scorer for multiclass classification" params.scorer = self.scorer @@ -442,4 +445,4 @@ def __setstate__(self, state): # self.data_ = None # self.train_ = None # self.validation_ = None - # self.search_space_ = None \ No newline at end of file + # self.search_space_ = None diff --git a/pybrush/deap_api/DeapEstimator.py b/pybrush/deap_api/DeapEstimator.py index ee7dde87b..176defe1a 100644 --- a/pybrush/deap_api/DeapEstimator.py +++ b/pybrush/deap_api/DeapEstimator.py @@ -156,6 +156,10 @@ def fit(self, X, y): "encoding method to convert the data to a supported " "format.") + if self.mode == 'classification': + self.classes_, y = np.unique(y, return_inverse=True) + y = y.astype(np.float32) + self.data_ = self._make_data(X, y, feature_names=self.feature_names_, feature_types=self.feature_types_, @@ -267,7 +271,10 @@ def predict(self, X): data = Dataset(X=X, ref_dataset=self.data_, feature_names=self.feature_names_) - return self.best_estimator_.program.predict(data) + prediction = np.asarray(self.best_estimator_.program.predict(data)) + if self.mode == 'classification': + return self.classes_[prediction.astype(int)] + return prediction # def _setup_population(self): # """initialize programs""" @@ -411,4 +418,4 @@ def __init__(self, **kwargs): # def transform(self, X): # """Transform X using the best estimator in the archive. """ -# return self.predict(X) \ No newline at end of file +# return self.predict(X) diff --git a/src/bindings/bind_engines.h b/src/bindings/bind_engines.h index ff788d22f..544521a03 100644 --- a/src/bindings/bind_engines.h +++ b/src/bindings/bind_engines.h @@ -89,14 +89,15 @@ void bind_engine(py::module& m, string name) ; // specialization for subclasses - if constexpr (std::is_same_v) + if constexpr (std::is_same_v || std::is_same_v) { + using ProbType = std::conditional_t, ArrayXXf, ArrayXf>; engine.def("predict_proba", - static_cast(&T::predict_proba), + static_cast(&T::predict_proba), "predict from Dataset object") .def("predict_proba", - static_cast &X)>(&T::predict_proba), + static_cast &X)>(&T::predict_proba), "predict from X data") ; } -} \ No newline at end of file +} diff --git a/src/bindings/bind_individuals.h b/src/bindings/bind_individuals.h index dc40b3e5d..863731737 100644 --- a/src/bindings/bind_individuals.h +++ b/src/bindings/bind_individuals.h @@ -79,15 +79,16 @@ void bind_individual(py::module& m, string name) ) ; - if constexpr (std::is_same_v) + if constexpr (std::is_same_v || std::is_same_v) { + using ProbType = typename br::Program::TreeType; ind.def("predict_proba", - static_cast(&Class::predict_proba), + static_cast(&Class::predict_proba), "predict from Dataset object") .def("predict_proba", - static_cast &X)>(&Class::predict_proba), + static_cast &X)>(&Class::predict_proba), "predict from X data") ; } -} \ No newline at end of file +} diff --git a/src/bindings/bind_programs.h b/src/bindings/bind_programs.h index e9464cec3..27f7af4c4 100644 --- a/src/bindings/bind_programs.h +++ b/src/bindings/bind_programs.h @@ -85,15 +85,16 @@ void bind_program(py::module& m, string name) ) ) ; - if constexpr (std::is_same_v) + if constexpr (std::is_same_v || std::is_same_v) { + using ProbType = typename T::TreeType; prog.def("predict_proba", - static_cast(&T::predict_proba), + static_cast(&T::predict_proba), "predict from Dataset object") .def("predict_proba", - static_cast &X)>(&T::predict_proba), + static_cast &X)>(&T::predict_proba), "predict from X data") ; } -} \ No newline at end of file +} diff --git a/src/eval/evaluation.cpp b/src/eval/evaluation.cpp index 0f2442204..b3bfdadf6 100644 --- a/src/eval/evaluation.cpp +++ b/src/eval/evaluation.cpp @@ -42,7 +42,7 @@ void Evaluation::update_fitness(Population& pop, // assign weights to individual if (fit && ind.get_is_fitted() == false) { - ind.program.fit(data.get_training_data()); + ind.program.fit(data.get_training_data(), params.class_weights); } assign_fit(ind, data, params, validation); @@ -126,4 +126,4 @@ void Evaluation::assign_fit(Individual& ind, const Dataset& data, template class Brush::Eval::Evaluation; template class Brush::Eval::Evaluation; template class Brush::Eval::Evaluation; -template class Brush::Eval::Evaluation; \ No newline at end of file +template class Brush::Eval::Evaluation; diff --git a/src/eval/metrics.cpp b/src/eval/metrics.cpp index 16a3f8549..79cb4c630 100644 --- a/src/eval/metrics.cpp +++ b/src/eval/metrics.cpp @@ -1,5 +1,7 @@ #include "metrics.h" +#include + namespace Brush { namespace Eval { @@ -234,51 +236,21 @@ float average_precision_score(const VectorXf& y, const VectorXf& predict_proba, VectorXf multi_log_loss(const VectorXf& y, const ArrayXXf& predict_proba, const vector& class_weights) { - // TODO: fix softmax and multiclassification, then implement this - VectorXf loss = VectorXf::Zero(y.rows()); - - // TODO: needs to be the index of unique elements - // get class labels - // vector uc = unique( ArrayXi(y.cast()) ); - - // float eps = 1e-6f; - // float sum_weights = 0; - // for (unsigned i = 0; i < y.rows(); ++i) - // { - // for (const auto& c : uc) - // { - // // for specific class - // ArrayXf yhat = predict_proba.col(int(c)); - - - // /* float yi = y(i) == c ? 1.0 : 0.0 ; */ - - // if (y(i) == c) - // { - // if (yhat(i) < eps || 1 - yhat(i) < eps) - // { - // // clip probabilities since log loss is undefined for yhat=0 or yhat=1 - // loss(i) += -log(eps); - // } - // else - // { - // loss(i) += -log(yhat(i)); - // } - - // } - - // } - // if (!class_weights.empty()){ - - // loss(i) = loss(i)*class_weights.at(y(i)); - // sum_weights += class_weights.at(y(i)); - // } - // } - // if (sum_weights > 0) - // loss = loss.array() / sum_weights * y.size(); - + if (predict_proba.rows() != y.rows()) + HANDLE_ERROR_THROW("Multiclass probabilities and labels have different numbers of rows"); + constexpr float eps = 1e-6f; + VectorXf loss(y.rows()); + for (int i = 0; i < y.rows(); ++i) + { + const int label = static_cast(y(i)); // labels are always encoded as integers for clf/multiclf + + // if (label < 0 || label >= predict_proba.cols()) + // HANDLE_ERROR_THROW("Class label is outside the predicted probability columns"); + // per sample log loss + loss(i) = -std::log(std::clamp(predict_proba(i, label), eps, 1.0f - eps)); + } return loss; } @@ -288,57 +260,84 @@ float mean_multi_log_loss(const VectorXf& y, { loss = multi_log_loss(y, predict_proba, class_weights); - return loss.mean(); + if (class_weights.empty()) + return loss.mean(); + + // apply class weights to the log loss + float sum_weights = 0.0f; + float weighted_loss = 0.0f; + for (int i = 0; i < y.rows(); ++i) + { + const float weight = class_weights.at(static_cast(y(i))); + weighted_loss += loss(i) * weight; + sum_weights += weight; + } + return sum_weights == 0.0f ? 0.0f : weighted_loss / sum_weights; } float multi_zero_one_loss(const VectorXf& y, const ArrayXXf& predict_proba, VectorXf& loss, const vector& class_weights ) { - // TODO: implement this - // vector uc = unique(y); - // vector c; - // for (const auto& i : uc) - // c.push_back(int(i)); - - // // sensitivity (TP) and specificity (TN) - // vector TP(c.size(),0.0), TN(c.size(), 0.0), P(c.size(),0.0), N(c.size(),0.0); - // ArrayXf class_accuracies(c.size()); - - // // get class counts - - // for (unsigned i=0; i< c.size(); ++i) - // { - // P.at(i) = (y.array().cast() == c.at(i)).count(); // total positives for this class - // N.at(i) = (y.array().cast() != c.at(i)).count(); // total negatives for this class - // } - + if (predict_proba.rows() != y.rows()) + HANDLE_ERROR_THROW("Multiclass probabilities and labels have different numbers of rows"); - // for (unsigned i = 0; i < y.rows(); ++i) - // { - // if (yhat(i) == y(i)) // true positive - // ++TP.at(y(i) == -1 ? 0 : y(i)); // if-then ? accounts for -1 class encoding + ArrayXi yhat(y.rows()); + for (int i = 0; i < predict_proba.rows(); ++i) + predict_proba.row(i).maxCoeff(&yhat(i)); // pick the predicted class - // for (unsigned j = 0; j < c.size(); ++j) - // if ( y(i) !=c.at(j) && yhat(i) != c.at(j) ) // true negative - // ++TN.at(j); - - // } + loss = (yhat.array() != y.cast().array()).cast(); // check if it was a hit or a miss - // // class-wise accuracy = 1/2 ( true positive rate + true negative rate) - // for (unsigned i=0; i< c.size(); ++i){ - // class_accuracies(i) = (TP.at(i)/P.at(i) + TN.at(i)/N.at(i))/2; + if (class_weights.empty()) // accuracy + return 1.0f - loss.mean(); + float weighted_errors = 0.0f; + float sum_weights = 0.0f; + for (int i = 0; i < y.rows(); ++i) + { + const float weight = class_weights.at(static_cast(y(i))); + weighted_errors += loss(i) * weight; + sum_weights += weight; + } + return sum_weights == 0.0f ? 0.0f : 1.0f - weighted_errors / sum_weights; +} +float multi_bal_zero_one_loss(const VectorXf& y, + const ArrayXXf& predict_proba, VectorXf& loss, + const vector& class_weights) +{ + if (predict_proba.rows() != y.rows()) + HANDLE_ERROR_THROW("Multiclass probabilities and labels have different numbers of rows"); - // } - - // // set loss vectors if third argument supplied - // loss = (yhat.cast().array() != y.cast().array()).cast(); + ArrayXi yhat(y.rows()); + for (int i = 0; i < predict_proba.rows(); ++i) + predict_proba.row(i).maxCoeff(&yhat(i)); + loss = (yhat.array() != y.cast().array()).cast(); - // return 1.0 - class_accuracies.mean(); - - return 0.0; + VectorXf correct = VectorXf::Zero(predict_proba.cols()); + VectorXf support = VectorXf::Zero(predict_proba.cols()); + for (int i = 0; i < y.rows(); ++i) + { + const int label = static_cast(y(i)); + + // if (label < 0 || label >= predict_proba.cols()) + // HANDLE_ERROR_THROW("Class label is outside the predicted probability columns"); + + // balanced, weighted by support + support(label) += 1.0f; + if (yhat(i) == label) + correct(label) += 1.0f; + } + + float recall_sum = 0.0f; + int present_classes = 0; + for (int label = 0; label < support.size(); ++label) + if (support(label) > 0.0f) + { + recall_sum += correct(label) / support(label); + ++present_classes; + } + return present_classes == 0 ? 0.0f : recall_sum / present_classes; } } // metrics diff --git a/src/eval/metrics.h b/src/eval/metrics.h index 3b42f69a0..1bafe7923 100644 --- a/src/eval/metrics.h +++ b/src/eval/metrics.h @@ -121,8 +121,13 @@ float multi_zero_one_loss(const VectorXf& y, const ArrayXXf& predict_proba, VectorXf& loss, const vector& class_weights=vector() ); +/// Balanced accuracy for multi-classification. +float multi_bal_zero_one_loss(const VectorXf& y, const ArrayXXf& predict_proba, + VectorXf& loss, + const vector& class_weights=vector() ); + } // metrics } // Brush -#endif \ No newline at end of file +#endif diff --git a/src/eval/scorer.h b/src/eval/scorer.h index a745fba97..19a4548c6 100644 --- a/src/eval/scorer.h +++ b/src/eval/scorer.h @@ -170,6 +170,7 @@ typedef float (*funcPointer)(const VectorXf&, Scorer(string scorer="multi_log") { score_hash["multi_log"] = &mean_multi_log_loss; score_hash["accuracy"] = &multi_zero_one_loss; + score_hash["balanced_accuracy"] = &multi_bal_zero_one_loss; this->set_scorer(scorer); }; diff --git a/src/program/functions.h b/src/program/functions.h index ababd344e..7019d7044 100644 --- a/src/program/functions.h +++ b/src/program/functions.h @@ -415,8 +415,14 @@ namespace Brush template inline auto softmax(const ArrayBase &t) const { - auto tMinusMax = t.rowwise() - t.colwise().maxCoeff(); - return tMinusMax.rowwise() - tMinusMax.exp().colwise().sum().log(); + // Rows are samples and columns are classes. Normalize each row, + // using the row maximum for numerical stability. + using Scalar = typename T::Scalar; + Array shifted = t; + shifted.colwise() -= shifted.rowwise().maxCoeff(); + auto exponentiated = shifted.exp(); + return (exponentiated.colwise() / + exponentiated.rowwise().sum()).eval(); } template @@ -425,12 +431,20 @@ namespace Brush return this->softmax(t); } - // template - // inline auto operator()(const Array& first, const Ts& ... inputs) - // { - // auto output = Stack(first, inputs...); - // return this->softmax(output); - // } + // N-ary Softmax receives one logits vector per class. + template + requires(sizeof...(Ts) > 0) + inline auto operator()(const ArrayBase& first, + const ArrayBase& ...inputs) + { + using Scalar = typename T::Scalar; + Array logits(first.rows(), + 1 + sizeof...(inputs)); + logits.col(0) = first; + int column = 1; + ((logits.col(column++) = inputs), ...); + return this->softmax(logits); + } }; /* logical and -- boolean AND operation */ diff --git a/src/program/optimizer/weight_optimizer.h b/src/program/optimizer/weight_optimizer.h index 0ff9a7fe6..5d92e3d07 100644 --- a/src/program/optimizer/weight_optimizer.h +++ b/src/program/optimizer/weight_optimizer.h @@ -30,11 +30,13 @@ struct OptimizerSummary { template struct ResidualEvaluator { typedef float Scalar; - ResidualEvaluator(PT& program, Dataset const& dataset) + ResidualEvaluator(PT& program, Dataset const& dataset, + const vector& class_weights = {}) : program_(program) , dataset_(dataset) , numParameters_(program.get_weights().size()) , y_true_(dataset.y) + , class_weights_(class_weights) {} template @@ -46,19 +48,34 @@ struct ResidualEvaluator { template auto operator()(T const* parameters, T* residuals) const -> bool { - using ArrayType = Array; // ColMajor? const T ** new_weights = ¶meters; + auto residualMap = Eigen::Map>( + residuals, GetDataset().get_n_samples()); - ArrayType y_pred = GetProgram().template predict_with_weights( - GetDataset(), - new_weights - ); - - auto residualMap = ArrayType::Map(residuals, GetDataset().get_n_samples()); - - // how we calculate the residuals - if (GetDataset().classification) // classification + if constexpr (PT::program_type == ProgramType::MulticlassClassifier) + { + using MatrixType = Array; + auto probabilities = GetProgram().template predict_with_weights( + GetDataset(), new_weights); + for (int i = 0; i < GetDataset().get_n_samples(); ++i) + { + const int label = static_cast(GetTarget()(i)); + const T probability = probabilities(i, label); + const float weight = class_weights_.empty() ? 1.0f + : class_weights_.at(label); + // TinySolver minimizes squared residuals, so sqrt(loss) + // makes its objective the weighted multinomial log loss. + residualMap(i) = sqrt(T(weight) * -log(probability)); + } + } + else { + using ArrayType = Array; + ArrayType y_pred = GetProgram().template predict_with_weights( + GetDataset(), new_weights); + + if (GetDataset().classification) + { // tolerance to avoid numeric errors. // Using an eps with 7 significant digits to avoid weird behavior. @@ -77,12 +94,21 @@ struct ResidualEvaluator { // clamp values and avoid log(0) y_pred = y_pred.min(T(1.0) - T(eps)).max(T(eps)); - // log loss - // residualMap = -(y*log(y_pred.array()) + (T(1.0)-y)*log(T(1.0)-y_pred.array())); - residualMap = -(y*log(y_pred) + (T(1.0)-y)*log(T(1.0)-y_pred)); - } - else { // This is MSE, default behavior - residualMap = (y_pred - GetTarget()); + for (int i = 0; i < y_pred.size(); ++i) + { + const float weight = class_weights_.empty() ? 1.0f + : class_weights_.at(static_cast(y(i))); + const T log_loss = -(T(y(i)) * log(y_pred(i)) + + (T(1.0f) - T(y(i))) * log(T(1.0f) - y_pred(i))); + // See multiclass branch above: this makes the least-squares + // objective equal weighted binary log loss. + residualMap(i) = sqrt(T(weight) * log_loss); + } + } + else + { + residualMap = y_pred - GetTarget(); + } } return true; @@ -98,6 +124,7 @@ struct ResidualEvaluator { std::reference_wrapper program_; std::reference_wrapper dataset_; std::reference_wrapper y_true_; + vector class_weights_; size_t numParameters_; // cache the number of parameters in the tree }; @@ -109,7 +136,8 @@ struct WeightOptimizer /// @param program the program /// @param dataset the dataset template - void update(PT& program, const Dataset& dataset) + void update(PT& program, const Dataset& dataset, + const vector& class_weights = {}) { if (program.get_n_weights() == 0) return; @@ -118,7 +146,7 @@ struct WeightOptimizer auto init_weights = program.get_weights(); using CFType = Brush::TinyCostFunction> ; - ResidualEvaluator evaluator(program, dataset); + ResidualEvaluator evaluator(program, dataset, class_weights); CFType cost_function(evaluator); ceres::TinySolver solver; solver.options.max_num_iterations = 10; @@ -154,4 +182,4 @@ struct WeightOptimizer }; } // namespace Brush -#endif \ No newline at end of file +#endif diff --git a/src/program/program.h b/src/program/program.h index 8986211bc..73ee6c763 100644 --- a/src/program/program.h +++ b/src/program/program.h @@ -148,10 +148,17 @@ template struct Program } Program& fit(const Dataset& d) + { + return fit(d, {}); + } + + /// Fit the program with optional per-class weights for classification. + Program& fit(const Dataset& d, + const vector& class_weights) { TreeType out = Tree.begin().node->fit(d); this->is_fitted_ = true; - update_weights(d); + update_weights(d, class_weights); // this->valid = true; return *this; }; @@ -294,7 +301,8 @@ template struct Program * * @param d the dataset */ - void update_weights(const Dataset& d); + void update_weights(const Dataset& d, + const vector& class_weights); /// @brief returns the number of weights in the program. int get_n_weights() const @@ -316,7 +324,8 @@ template struct Program // It is important that this condition also matches the condition in // the methods get_weights and set_weights. if (Is(node.node_type) - || (node.get_is_weighted() && IsWeighable(node.ret_type)) ) + || (node.get_is_weighted() && IsWeighable(node.ret_type) + && IsWeighable(node.node_type)) ) ++count; } return count; @@ -340,7 +349,8 @@ template struct Program continue; if ( Is(node.node_type) - || (node.get_is_weighted() && IsWeighable(node.ret_type)) ) + || (node.get_is_weighted() && IsWeighable(node.ret_type) + && IsWeighable(node.node_type)) ) { weights(i) = node.W; ++i; @@ -372,7 +382,8 @@ template struct Program continue; if ( Is(node.node_type) - || (node.get_is_weighted() && IsWeighable(node.node_type)) ) + || (node.get_is_weighted() && IsWeighable(node.ret_type) + && IsWeighable(node.node_type)) ) { node.W = weights(j); ++j; @@ -601,6 +612,10 @@ template struct Program head_label = edge_label; } + else if (Is(parent->data.node_type)){ + // Each child supplies logits for one output class. + edge_label = fmt::format("class {}", j); + } // drawing the edges string font_color = ""; @@ -671,13 +686,14 @@ template struct Program namespace Brush{ template -void Program::update_weights(const Dataset& d) +void Program::update_weights(const Dataset& d, + const vector& class_weights) { // Updates the weights within a tree. // make an optimizer auto WO = WeightOptimizer(); // get new weights from optimization. - WO.update((*this), d); + WO.update((*this), d, class_weights); }; diff --git a/src/program/split.cpp b/src/program/split.cpp index b66268204..64ddd59ae 100644 --- a/src/program/split.cpp +++ b/src/program/split.cpp @@ -82,11 +82,13 @@ float gain(const ArrayXf& lsplit, } else { - lscore = variance(lsplit)/float(lsplit.size()); - rscore = variance(rsplit)/float(rsplit.size()); + lscore = variance(lsplit); + rscore = variance(rsplit); /* cout << "lscore: " << lscore << "\n"; */ /* cout << "rscore: " << rscore << "\n"; */ - score = lscore + rscore; + score = (lscore * float(lsplit.size()) + + rscore * float(rsplit.size())) + / float(lsplit.size() + rsplit.size()); } return score; @@ -108,4 +110,4 @@ float gini_impurity_index(const ArrayXf& classes, return gini; } -} //Brush::Split \ No newline at end of file +} //Brush::Split diff --git a/src/simplification/constants.h b/src/simplification/constants.h index 845e3b669..f58777e96 100644 --- a/src/simplification/constants.h +++ b/src/simplification/constants.h @@ -50,9 +50,21 @@ namespace Brush { namespace Simpl{ } else if constexpr (P==ProgramType::MulticlassClassifier) { - ArrayXXf out = (*spot.node).template predict(d); - auto argmax = Function{}; - branch_pred = ArrayXf(argmax(out).template cast()); + // A multiclass tree contains ArrayF logit branches + // below its MatrixF Softmax root. Evaluate each + // node using its own return type; treating a logit + // branch as MatrixF requests an invalid dispatch + // callable (for example Sum(MatrixF)). + if (n.ret_type == DataType::MatrixF) + { + ArrayXXf out = (*spot.node).template predict(d); + auto argmax = Function{}; + branch_pred = ArrayXf(argmax(out).template cast()); + } + else + { + branch_pred = (*spot.node).template predict(d); + } } else { diff --git a/src/vary/search_space.h b/src/vary/search_space.h index 75322dcb7..7ac85c38b 100644 --- a/src/vary/search_space.h +++ b/src/vary/search_space.h @@ -494,6 +494,37 @@ struct SearchSpace weights.end())); }; + /// @brief Get a specific operator with an exact number of arguments. + std::optional sample_op(NodeType type, DataType R, + size_t arg_count, + bool force_return=false) const + { + check(R); + if (node_map.find(R) == node_map.end()) + return std::nullopt; + + vector matches; + vector weights; + for (const auto& [arg_hash, node_type_map] : node_map.at(R)) + { + auto node = node_type_map.find(type); + if (node != node_type_map.end() + && node->second.arg_types.size() == arg_count) + { + matches.push_back(node->second); + weights.push_back(node_map_weights.at(R).at(arg_hash).at(type)); + } + } + + if (force_return) + std::fill(weights.begin(), weights.end(), 1.0f); + if (matches.empty() || !has_solution_space(weights.begin(), weights.end())) + return std::nullopt; + + return *r.select_randomly(matches.begin(), matches.end(), + weights.begin(), weights.end()); + }; + /// @brief get operator with at least one argument matching arg /// @param ret return type /// @param arg argument type to match @@ -761,7 +792,15 @@ P SearchSpace::make_program(const Parameters& params, int max_d, int max_size) } else if (P::program_type == ProgramType::MulticlassClassifier) { - Node node_softmax = sample_op(NodeType::Softmax, DataType::MatrixF, true).value(); + auto softmax = sample_op(NodeType::Softmax, DataType::MatrixF, + params.n_classes, true); + + if (!softmax) // should never happen. Let's keep this here just so we can catch if it ever happens + HANDLE_ERROR_THROW(fmt::format( + "Multiclass Softmax supports between 2 and {} classes; got {}.\n", + MAX_ARGS, params.n_classes)); + + Node node_softmax = softmax.value(); node_softmax.set_prob_change(0.0); node_softmax.set_is_weighted(false); // same as logistic roots diff --git a/src/vary/variation.h b/src/vary/variation.h index 02a2b06f8..4917fdcf6 100644 --- a/src/vary/variation.h +++ b/src/vary/variation.h @@ -343,7 +343,7 @@ class Variation { assert(ind.program.size() > 0); assert(ind.fitness.valid() == false); - ind.program.fit(data.get_training_data()); + ind.program.fit(data.get_training_data(), parameters.class_weights); // simplify before calculating fitness (order matters, as they are not refitted and constants simplifier does not replace with the right value.) // simplify constants first to avoid letting the lsh simplifier to visit redundant branches @@ -662,4 +662,4 @@ extern template class Variation; } //namespace Var } //namespace Brush -#endif \ No newline at end of file +#endif diff --git a/tests/cpp/test_evaluation.cpp b/tests/cpp/test_evaluation.cpp index 9b40a462b..62ec48a8f 100644 --- a/tests/cpp/test_evaluation.cpp +++ b/tests/cpp/test_evaluation.cpp @@ -2,6 +2,8 @@ #include "../../src/eval/evaluation.h" #include "../../src/eval/metrics.h" #include "../../src/eval/scorer.h" +#include "../../src/program/functions.h" +#include "../../src/simplification/constants.h" using namespace Brush::Eval; @@ -84,6 +86,136 @@ TEST(Evaluation, ScorerBinaryAccuracy) ASSERT_TRUE(loss.isApprox(loss_expected, 1e-6)); } +TEST(Evaluation, MulticlassSoftmaxAndMetrics) +{ + ArrayXf first(2), second(2), third(2); + first << 2.0f, 0.0f; + second << 1.0f, 1.0f; + third << 0.0f, 2.0f; + + // testing wether softmax return valid probabilities (normalized per row). + // softmax here is representing 3 classes, and we have 2 rows (2 samples) + const auto probabilities = Function{}(first, second, third); + + ASSERT_EQ(probabilities.rows(), 2); + ASSERT_EQ(probabilities.cols(), 3); + + ASSERT_NEAR(probabilities.row(0).sum(), 1.0f, 1e-6f); + ASSERT_NEAR(probabilities.row(1).sum(), 1.0f, 1e-6f); + + EXPECT_GT(probabilities(0, 0), probabilities(0, 1)); + EXPECT_GT(probabilities(1, 2), probabilities(1, 1)); + + VectorXf y(2), loss(2); + y << 0.0f, 2.0f; // y needs to be float, but gets casted to int when checking for hit/miss inside multiclass losses + // The model has 100% accuracy for this y, so we should expect perfect metric values + EXPECT_NEAR(mean_multi_log_loss(y, probabilities, loss), + (-std::log(probabilities(0, 0)) - std::log(probabilities(1, 2))) / 2.0f, + 1e-6f); + EXPECT_NEAR(multi_zero_one_loss(y, probabilities, loss), 1.0f, 1e-6f); + EXPECT_NEAR(multi_bal_zero_one_loss(y, probabilities, loss), 1.0f, 1e-6f); +} + +TEST(Evaluation, MulticlassSoftmaxHasOneOutputPerClass) +{ + ArrayXXf X(6, 2); + X << 0.0f, 1.0f, + 1.0f, 0.0f, + 0.5f, 0.5f, + 2.0f, 1.0f, + 1.0f, 2.0f, + 2.0f, 2.0f; + ArrayXf y(6); + y << 0.0f, 1.0f, 2.0f, 0.0f, 1.0f, 2.0f; + + Dataset data(X, y, {}, {}, {}, true); + SearchSpace search_space(data); + Parameters params; + params.classification = true; + params.set_n_classes(y); + params.max_depth = 3; + params.max_size = 20; + + auto program = search_space.make_multiclass_classifier(3, 20, params); + + // root softmax must have one arg (subtree) for each class + EXPECT_EQ(program.Tree.begin().node->data.arg_types.size(), params.n_classes); + + // something needs to be tunable. Even without constants, we have the softmax weights + ASSERT_GT(program.get_n_weights(), 0); + + program.fit(data, {1.0f, 2.0f, 3.0f}); + + const auto probabilities = program.predict_proba(data); + EXPECT_EQ(probabilities.cols(), params.n_classes); + + // each row probability sums up to 1 + for (int row = 0; row < probabilities.rows(); ++row) + EXPECT_NEAR(probabilities.row(row).sum(), 1.0f, 1e-6f); + + Simpl::Constants_simplifier simplifier; + simplifier.simplify_tree(program, search_space, data); + + // simplification, if applied, does not break softmax + const auto simplified_probabilities = program.predict_proba(data); + for (int row = 0; row < simplified_probabilities.rows(); ++row) + EXPECT_NEAR(simplified_probabilities.row(row).sum(), 1.0f, 1e-6f); +} + +TEST(Evaluation, SplitThresholdsAreNotOptimizerParameters) +{ + ArrayXXf X(6, 1); + X << 0.0f, 1.0f, 2.0f, 3.0f, 4.0f, 5.0f; + + ArrayXf y(6); + y << 0.0f, 0.0f, 0.0f, 1.0f, 1.0f, 1.0f; + + Dataset data(X, y, {}, {}, {"ArrayF"}); + + RegressorProgram program = json({{"Tree", { + {{"node_type", "SplitBest"}, {"is_weighted", true}}, + {{"node_type", "Constant"}, {"is_weighted", true}}, + {{"node_type", "Constant"}, {"is_weighted", true}} + }}, {"is_fitted_", false}}); + + program.fit(data); + const auto prediction = program.predict(data); + + EXPECT_TRUE(prediction.isApprox(y, 1e-3f)) + << "prediction=" << prediction.transpose() + << ", weights=" << program.get_weights().transpose() + << ", model=" << program.get_model(); + + // The greedy split threshold lives in SplitBest::W and must not enter the + // Ceres parameter vector; only the two leaf constants do. + EXPECT_EQ(program.get_weights().size(), 2); // split best should not be considered, so we will have 2 optimizable parameters on this tree +} + +TEST(Evaluation, BinaryOptimizerUsesClassWeights) +{ + ArrayXXf X(2, 1); + X << 0.0f, 1.0f; + + ArrayXf y(2); + y << 0.0f, 1.0f; + + Dataset data(X, y, {}, {}, {}, true); + + // given the X and y above, we should have a perfect classifier here + ClassifierProgram program = json({{"Tree", { + {{"node_type", "Logistic"}, {"is_weighted", false}}, + {{"node_type", "OffsetSum"}, {"is_weighted", true}}, + {{"node_type", "Constant"}, {"is_weighted", true}} + }}, {"is_fitted_", false}}); + + program.fit(data, {1.0f, 9.0f}); + + const auto probability = program.predict_proba(data); + + // one proba will be low, the other will be super high + EXPECT_GT(probability.mean(), 0.8f); +} + // TEST(EvaluationTest, UpdateFitnessTest) { // // TODO: Add test case for update_fitness function diff --git a/tests/python/test_deap_api.py b/tests/python/test_deap_api.py index aed89ff90..318f13f81 100644 --- a/tests/python/test_deap_api.py +++ b/tests/python/test_deap_api.py @@ -5,6 +5,7 @@ import pandas as pd from pmlb import fetch_data from sklearn.utils import resample +from sklearn.datasets import make_classification import traceback import logging @@ -149,6 +150,25 @@ def test_predict_proba(setup, brush_args, request): "every class should have its own column (even for binary clf)" +def test_deap_multiclass_probabilities_and_labels(): + X, y = make_classification( + n_samples=36, n_features=4, n_informative=3, n_redundant=0, + n_classes=3, n_clusters_per_class=1, random_state=12) + labels = np.array(["red", "green", "blue"])[y] + + est = pybrush.deap_api.DeapClassifier( + max_gens=2, pop_size=8, max_size=30, max_depth=4, + num_islands=1, validation_size=0.0, random_state=12) + est.fit(X, labels) + + probabilities = est.predict_proba(X) + prediction = est.predict(X) + assert probabilities.shape == (X.shape[0], 3) + assert np.allclose(probabilities.sum(axis=1), 1.0) + assert set(prediction).issubset(set(labels)) + assert np.array_equal(est.classes_, np.array(["blue", "green", "red"])) + + # @pytest.mark.parametrize('setup,num_islands', # [('DEAP_classification_setup', 1), # ('DEAP_regression_setup', 1), @@ -247,4 +267,4 @@ def test_fixed_nodes(setup, fixed_node, brush_args, request): # est2 = pybrush.BrushRegressor(random_state=42).fit(test_X, test_y) # assert est1.best_estimator_.program.get_model() == est2.best_estimator_.program.get_model(), \ -# "random state failed to generate same results" \ No newline at end of file +# "random state failed to generate same results" diff --git a/tests/python/test_sklearn_interface.py b/tests/python/test_sklearn_interface.py index b1bf949df..6965b63a4 100644 --- a/tests/python/test_sklearn_interface.py +++ b/tests/python/test_sklearn_interface.py @@ -204,5 +204,24 @@ def test_brush_lock_nodes_and_leaves(): assert fitness_after.loss_v >= fitness_before.loss_v +def test_brush_multiclass_probabilities_and_labels(): + X, y = make_classification( + n_samples=45, n_features=5, n_informative=4, n_redundant=0, + n_classes=3, n_clusters_per_class=1, random_state=11) + labels = np.array(["class-a", "class-b", "class-c"])[y] + + est = BrushClassifier( + max_gens=2, pop_size=8, max_size=30, max_depth=4, + num_islands=1, validation_size=0.0, random_state=11) + est.fit(X, labels) + + probabilities = est.predict_proba(X) + prediction = est.predict(X) + assert probabilities.shape == (X.shape[0], 3) + assert np.allclose(probabilities.sum(axis=1), 1.0) + assert set(prediction).issubset(set(labels)) + assert np.array_equal(est.classes_, np.array(["class-a", "class-b", "class-c"])) + + if __name__ == "__main__": - pytest.main() \ No newline at end of file + pytest.main()