11
Datasets
ML
Tensors
Metrics
Tree-Based & Ensemble Models
Decision Trees, Random Forests, Gradient Boosting, and Linear SVM. Covers both classification and regression variants. This example uses deepbox/datasets, deepbox/ml, deepbox/ndarray, deepbox/metrics, deepbox/preprocess and focuses on loadIris; DecisionTree, RandomForest, GradientBoosting (Classifier + Regressor), LinearSVC, LinearSVR; tensor, slice; accuracy, mse, r2Score; trainTestSplit.
Deepbox Modules Used
deepbox/datasetsdeepbox/mldeepbox/ndarraydeepbox/metricsdeepbox/preprocessWhat You Will Learn
- Use deepbox/datasets for loadIris.
- Use deepbox/ml for DecisionTree, RandomForest, GradientBoosting (Classifier + Regressor), LinearSVC, LinearSVR.
- Use deepbox/ndarray for tensor, slice.
- Use deepbox/metrics for accuracy, mse, r2Score.
- Decision Trees, Random Forests, Gradient Boosting, and Linear SVM. Covers both classification and regression variants.
Source Files
index.ts
1/**2 * Example 11: Tree-Based & Ensemble Models3 *4 * Decision Trees, Random Forests, Gradient Boosting, and Linear SVM.5 * Covers classification and regression variants.6 */78import { loadIris } from "deepbox/datasets";9import { accuracy, mse, r2Score } from "deepbox/metrics";10import {11 DecisionTreeClassifier,12 DecisionTreeRegressor,13 GradientBoostingClassifier,14 GradientBoostingRegressor,15 LinearSVC,16 LinearSVR,17 RandomForestClassifier,18 RandomForestRegressor,19} from "deepbox/ml";20import { slice, tensor } from "deepbox/ndarray";21import { trainTestSplit } from "deepbox/preprocess";2223console.log("=== Tree-Based & Ensemble Models ===\n");2425// ---------------------------------------------------------------------------26// Classification dataset (Iris — full 3-class for multi-class models)27// ---------------------------------------------------------------------------28const iris = loadIris();29const [XTrain, XTest, yTrain, yTest] = trainTestSplit(iris.data, iris.target, {30 testSize: 0.2,31 randomState: 42,32});3334// Binary subset (classes 0 and 1 only) for models that require binary labels35const XBin = slice(iris.data, { start: 0, end: 100 });36const yBin = slice(iris.target, { start: 0, end: 100 });37const [XBinTrain, XBinTest, yBinTrain, yBinTest] = trainTestSplit(XBin, yBin, {38 testSize: 0.2,39 randomState: 42,40});4142// ---------------------------------------------------------------------------43// Part 1: Decision Tree Classifier44// ---------------------------------------------------------------------------45console.log("--- Part 1: Decision Tree Classifier ---");4647const dtc = new DecisionTreeClassifier({ maxDepth: 5, minSamplesSplit: 2 });48dtc.fit(XTrain, yTrain);49const dtcPred = dtc.predict(XTest);50console.log(" Accuracy:", accuracy(yTest, dtcPred).toFixed(4));5152// ---------------------------------------------------------------------------53// Part 2: Random Forest Classifier54// ---------------------------------------------------------------------------55console.log("\n--- Part 2: Random Forest Classifier ---");5657const rfc = new RandomForestClassifier({58 nEstimators: 50,59 maxDepth: 5,60 randomState: 42,61});62rfc.fit(XTrain, yTrain);63const rfcPred = rfc.predict(XTest);64console.log(" Accuracy:", accuracy(yTest, rfcPred).toFixed(4));6566// ---------------------------------------------------------------------------67// Part 3: Gradient Boosting Classifier68// ---------------------------------------------------------------------------69console.log("\n--- Part 3: Gradient Boosting Classifier ---");7071const gbc = new GradientBoostingClassifier({72 nEstimators: 50,73 learningRate: 0.1,74 maxDepth: 3,75});76gbc.fit(XBinTrain, yBinTrain);77const gbcPred = gbc.predict(XBinTest);78console.log(" Accuracy:", accuracy(yBinTest, gbcPred).toFixed(4));7980// ---------------------------------------------------------------------------81// Part 4: Linear SVC82// ---------------------------------------------------------------------------83console.log("\n--- Part 4: Linear SVC ---");8485const svc = new LinearSVC({ C: 1.0 });86svc.fit(XBinTrain, yBinTrain);87const svcPred = svc.predict(XBinTest);88console.log(" Accuracy:", accuracy(yBinTest, svcPred).toFixed(4));8990// ---------------------------------------------------------------------------91// Regression dataset (synthetic y = x0 + 2*x1 + noise)92// ---------------------------------------------------------------------------93console.log("\n--- Regression Models ---");9495const XReg = tensor([96 [1, 2],97 [2, 3],98 [3, 4],99 [4, 5],100 [5, 6],101 [6, 7],102 [7, 8],103 [8, 9],104 [9, 10],105 [10, 11],106 [1, 3],107 [2, 5],108 [3, 2],109 [4, 1],110 [5, 4],111 [6, 3],112 [7, 6],113 [8, 5],114 [9, 8],115 [10, 7],116]);117const yReg = tensor([5, 8, 11, 14, 17, 20, 23, 26, 29, 32, 7, 12, 7, 6, 13, 12, 19, 18, 25, 24]);118119const [XRegTrain, XRegTest, yRegTrain, yRegTest] = trainTestSplit(XReg, yReg, {120 testSize: 0.2,121 randomState: 42,122});123124// ---------------------------------------------------------------------------125// Part 5: Decision Tree Regressor126// ---------------------------------------------------------------------------127console.log("\n--- Part 5: Decision Tree Regressor ---");128129const dtr = new DecisionTreeRegressor({ maxDepth: 5 });130dtr.fit(XRegTrain, yRegTrain);131const dtrPred = dtr.predict(XRegTest);132console.log(" MSE:", mse(yRegTest, dtrPred).toFixed(4));133console.log(" R²: ", r2Score(yRegTest, dtrPred).toFixed(4));134135// ---------------------------------------------------------------------------136// Part 6: Random Forest Regressor137// ---------------------------------------------------------------------------138console.log("\n--- Part 6: Random Forest Regressor ---");139140const rfr = new RandomForestRegressor({141 nEstimators: 50,142 maxDepth: 5,143 randomState: 42,144});145rfr.fit(XRegTrain, yRegTrain);146const rfrPred = rfr.predict(XRegTest);147console.log(" MSE:", mse(yRegTest, rfrPred).toFixed(4));148console.log(" R²: ", r2Score(yRegTest, rfrPred).toFixed(4));149150// ---------------------------------------------------------------------------151// Part 7: Gradient Boosting Regressor152// ---------------------------------------------------------------------------153console.log("\n--- Part 7: Gradient Boosting Regressor ---");154155const gbr = new GradientBoostingRegressor({156 nEstimators: 50,157 learningRate: 0.1,158 maxDepth: 3,159});160gbr.fit(XRegTrain, yRegTrain);161const gbrPred = gbr.predict(XRegTest);162console.log(" MSE:", mse(yRegTest, gbrPred).toFixed(4));163console.log(" R²: ", r2Score(yRegTest, gbrPred).toFixed(4));164165// ---------------------------------------------------------------------------166// Part 8: Linear SVR167// ---------------------------------------------------------------------------168console.log("\n--- Part 8: Linear SVR ---");169170const svr = new LinearSVR({ C: 1.0 });171svr.fit(XRegTrain, yRegTrain);172const svrPred = svr.predict(XRegTest);173console.log(" MSE:", mse(yRegTest, svrPred).toFixed(4));174console.log(" R²: ", r2Score(yRegTest, svrPred).toFixed(4));175176console.log("\n=== Tree-Based & Ensemble Models Complete ===");177Console Output
$ npx tsx 11-tree-ensemble-models/index.ts
Console output showing accuracy and regression metrics for 8 different tree-based and ensemble models