Example 11
intermediate
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/preprocess

What 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 ===");177

Console Output

$ npx tsx 11-tree-ensemble-models/index.ts
Console output showing accuracy and regression metrics for 8 different tree-based and ensemble models