diff --git a/ensemble/src/main/java/org/sklearn/ensemble/AdaBoostClassifier.java b/ensemble/src/main/java/org/sklearn/ensemble/AdaBoostClassifier.java
new file mode 100644
index 0000000..3be66d4
--- /dev/null
+++ b/ensemble/src/main/java/org/sklearn/ensemble/AdaBoostClassifier.java
@@ -0,0 +1,370 @@
+package org.sklearn.ensemble;
+
+import org.sklearn.core.Predictor;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.RandomGenerator;
+import org.sklearn.math.Vector;
+import org.sklearn.tree.DecisionTreeClassifier;
+import org.sklearn.utils.Validation;
+
+import java.util.*;
+
+/**
+ * AdaBoost classifier.
+ *
+ *
Fits a sequence of weak learners (default: decision stumps with
+ * {@code maxDepth=1}) on weighted versions of the training data. At each
+ * iteration sample weights are updated to give more weight to misclassified
+ * samples. Prediction uses weighted majority voting (SAMME algorithm).
+ *
+ *
Mirrors {@code sklearn.ensemble.AdaBoostClassifier}.
+ *
+ *
Usage:
+ *
{@code
+ * AdaBoostClassifier ada = new AdaBoostClassifier(50, 1, 42);
+ * ada.fit(X, y);
+ * Vector preds = ada.predict(X_test);
+ * }
+ */
+public class AdaBoostClassifier implements Predictor {
+
+ private int nEstimators;
+ private int maxDepth;
+ private double learningRate;
+ private long seed;
+ private boolean fitted;
+ private List estimators;
+ private double[] estimatorWeights;
+ private int[] classes;
+ private int nFeatures;
+
+ /**
+ * Create an AdaBoost classifier.
+ *
+ * @param nEstimators number of weak learners
+ * @param maxDepth max depth of each weak learner (default 1 = stump)
+ * @param seed random seed
+ */
+ public AdaBoostClassifier(int nEstimators, int maxDepth, long seed) {
+ this(nEstimators, maxDepth, 1.0, seed);
+ }
+
+ /**
+ * Create an AdaBoost classifier with full control.
+ *
+ * @param nEstimators number of weak learners
+ * @param maxDepth max depth of each weak learner
+ * @param learningRate learning rate (shrinks contribution of each learner)
+ * @param seed random seed
+ */
+ public AdaBoostClassifier(int nEstimators, int maxDepth,
+ double learningRate, long seed) {
+ if (nEstimators < 1) {
+ throw new IllegalArgumentException("nEstimators must be >= 1");
+ }
+ if (maxDepth < 1) {
+ throw new IllegalArgumentException("maxDepth must be >= 1");
+ }
+ if (learningRate <= 0) {
+ throw new IllegalArgumentException("learningRate must be > 0");
+ }
+ this.nEstimators = nEstimators;
+ this.maxDepth = maxDepth;
+ this.learningRate = learningRate;
+ this.seed = seed;
+ }
+
+ @Override
+ public AdaBoostClassifier fit(Matrix X, Vector y) {
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ int n = X.rows();
+ int m = X.cols();
+ this.nFeatures = m;
+
+ Set uniqueLabels = new LinkedHashSet<>();
+ for (int i = 0; i < n; i++) {
+ uniqueLabels.add((int) y.get(i));
+ }
+ this.classes = uniqueLabels.stream().mapToInt(Integer::intValue).toArray();
+ Arrays.sort(this.classes);
+ int nClasses = classes.length;
+
+ double[] sampleWeights = new double[n];
+ Arrays.fill(sampleWeights, 1.0 / n);
+
+ estimators = new ArrayList<>();
+ estimatorWeights = new double[nEstimators];
+ double epsilon = 1e-10;
+
+ for (int t = 0; t < nEstimators; t++) {
+ for (int i = 0; i < n; i++) {
+ if (sampleWeights[i] == 0) {
+ sampleWeights[i] = epsilon;
+ }
+ }
+
+ Object[] bootData = bootstrapSample(X, y, sampleWeights, t);
+ DecisionTreeClassifier stump = new DecisionTreeClassifier(
+ maxDepth, 2, 1, "gini");
+ stump.fit((Matrix) bootData[0], (Vector) bootData[1]);
+
+ Vector preds = stump.predict(X);
+ double error = 0.0;
+ for (int i = 0; i < n; i++) {
+ if (Math.abs(preds.get(i) - y.get(i)) >= 0.5) {
+ error += sampleWeights[i];
+ }
+ }
+ error /= sum(sampleWeights);
+
+ if (error > 1.0 - 1.0 / nClasses) {
+ estimatorWeights[t] = 0.0;
+ estimators.add(stump);
+ continue;
+ }
+
+ double alpha;
+ if (nClasses == 2) {
+ alpha = learningRate * 0.5 * Math.log((1 - error) / Math.max(error, epsilon));
+ } else {
+ alpha = learningRate * (Math.log((1 - error) / Math.max(error, epsilon))
+ + Math.log(nClasses - 1));
+ }
+
+ estimatorWeights[t] = alpha;
+ estimators.add(stump);
+
+ double sumW = 0.0;
+ for (int i = 0; i < n; i++) {
+ boolean misclassified = Math.abs(preds.get(i) - y.get(i)) >= 0.5;
+ sampleWeights[i] *= Math.exp(alpha * (misclassified ? 1 : -1));
+ sampleWeights[i] = Math.max(sampleWeights[i], epsilon);
+ sumW += sampleWeights[i];
+ }
+ for (int i = 0; i < n; i++) {
+ sampleWeights[i] /= sumW;
+ }
+
+ if (error == 0) {
+ break;
+ }
+ }
+
+ fitted = true;
+ return this;
+ }
+
+ private Object[] bootstrapSample(Matrix X, Vector y, double[] weights, long iterSeed) {
+ int n = X.rows();
+ int m = X.cols();
+ RandomGenerator rng = new RandomGenerator(seed + iterSeed * 1000);
+
+ double[] cumSum = new double[n];
+ cumSum[0] = weights[0];
+ for (int i = 1; i < n; i++) {
+ cumSum[i] = cumSum[i - 1] + weights[i];
+ }
+ double totalW = cumSum[n - 1];
+
+ Matrix bootX = new Matrix(n, m);
+ Vector bootY = new Vector(n);
+ for (int i = 0; i < n; i++) {
+ double r = rng.nextDouble() * totalW;
+ int idx = lowerBound(cumSum, r);
+ idx = Math.min(idx, n - 1);
+ for (int j = 0; j < m; j++) {
+ bootX.set(i, j, X.get(idx, j));
+ }
+ bootY.set(i, y.get(idx));
+ }
+ return new Object[]{bootX, bootY};
+ }
+
+ private int lowerBound(double[] arr, double val) {
+ int lo = 0, hi = arr.length;
+ while (lo < hi) {
+ int mid = (lo + hi) / 2;
+ if (arr[mid] < val) {
+ lo = mid + 1;
+ } else {
+ hi = mid;
+ }
+ }
+ return lo;
+ }
+
+ private double sum(double[] arr) {
+ double s = 0;
+ for (double v : arr) {
+ s += v;
+ }
+ return s;
+ }
+
+ @Override
+ public Vector predict(Matrix X) {
+ Validation.checkFitted(fitted, "AdaBoostClassifier");
+ Validation.checkMatrix(X, -1);
+ if (X.cols() != nFeatures) {
+ throw new IllegalArgumentException(
+ "Feature dimension mismatch: expected " + nFeatures
+ + " features, got " + X.cols());
+ }
+
+ int n = X.rows();
+ int nClasses = classes.length;
+
+ if (nClasses == 2) {
+ double[] weightedScores = new double[n];
+ for (int t = 0; t < estimators.size(); t++) {
+ if (estimatorWeights[t] == 0) {
+ continue;
+ }
+ Vector preds = estimators.get(t).predict(X);
+ for (int i = 0; i < n; i++) {
+ weightedScores[i] += estimatorWeights[t]
+ * (preds.get(i) == classes[1] ? 1 : -1);
+ }
+ }
+ double[] out = new double[n];
+ for (int i = 0; i < n; i++) {
+ out[i] = weightedScores[i] >= 0 ? classes[1] : classes[0];
+ }
+ return new Vector(out);
+ } else {
+ double[][] votes = new double[n][nClasses];
+ for (int t = 0; t < estimators.size(); t++) {
+ if (estimatorWeights[t] == 0) {
+ continue;
+ }
+ Vector preds = estimators.get(t).predict(X);
+ for (int i = 0; i < n; i++) {
+ int predClass = (int) preds.get(i);
+ int idx = classIndex(predClass);
+ votes[i][idx] += estimatorWeights[t];
+ }
+ }
+ double[] out = new double[n];
+ for (int i = 0; i < n; i++) {
+ out[i] = classes[argmax(votes[i])];
+ }
+ return new Vector(out);
+ }
+ }
+
+ /**
+ * Predict class probabilities based on normalized weighted votes.
+ */
+ public Vector predictProba(Matrix X) {
+ Validation.checkFitted(fitted, "AdaBoostClassifier");
+ Validation.checkMatrix(X, -1);
+
+ int n = X.rows();
+ int nClasses = classes.length;
+
+ double[] probs = new double[n];
+
+ if (nClasses == 2) {
+ double maxAbsWeight = 0;
+ double[] scores = new double[n];
+ for (int t = 0; t < estimators.size(); t++) {
+ if (estimatorWeights[t] == 0) {
+ continue;
+ }
+ Vector preds = estimators.get(t).predict(X);
+ double absW = Math.abs(estimatorWeights[t]);
+ maxAbsWeight += absW;
+ for (int i = 0; i < n; i++) {
+ scores[i] += estimatorWeights[t]
+ * (preds.get(i) == classes[1] ? 1 : -1);
+ }
+ }
+ for (int i = 0; i < n; i++) {
+ probs[i] = maxAbsWeight > 0
+ ? (scores[i] / maxAbsWeight + 1.0) / 2.0
+ : 0.5;
+ probs[i] = Math.max(0.0, Math.min(1.0, probs[i]));
+ }
+ } else {
+ double[][] votes = new double[n][nClasses];
+ for (int t = 0; t < estimators.size(); t++) {
+ if (estimatorWeights[t] == 0) {
+ continue;
+ }
+ Vector preds = estimators.get(t).predict(X);
+ for (int i = 0; i < n; i++) {
+ int idx = classIndex((int) preds.get(i));
+ votes[i][idx] += estimatorWeights[t];
+ }
+ }
+ for (int i = 0; i < n; i++) {
+ double minVote = Double.MAX_VALUE, maxVote = -Double.MAX_VALUE;
+ for (int k = 0; k < nClasses; k++) {
+ if (votes[i][k] < minVote) {
+ minVote = votes[i][k];
+ }
+ if (votes[i][k] > maxVote) {
+ maxVote = votes[i][k];
+ }
+ }
+ double range = maxVote - minVote;
+ if (range > 0) {
+ probs[i] = (votes[i][classIndex(classes[1])] - minVote) / range;
+ } else {
+ probs[i] = 0.5;
+ }
+ }
+ }
+ return new Vector(probs);
+ }
+
+ @Override
+ public double score(Matrix X, Vector y) {
+ Validation.checkFitted(fitted, "AdaBoostClassifier");
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ Vector pred = predict(X);
+ int correct = 0;
+ for (int i = 0; i < y.size(); i++) {
+ if (Math.abs(pred.get(i) - y.get(i)) < 0.5) {
+ correct++;
+ }
+ }
+ return (double) correct / y.size();
+ }
+
+ private int classIndex(int cls) {
+ for (int i = 0; i < classes.length; i++) {
+ if (classes[i] == cls) {
+ return i;
+ }
+ }
+ return 0;
+ }
+
+ private int argmax(double[] arr) {
+ int best = 0;
+ for (int i = 1; i < arr.length; i++) {
+ if (arr[i] > arr[best]) {
+ best = i;
+ }
+ }
+ return best;
+ }
+
+ public boolean isFitted() {
+ return fitted;
+ }
+
+ @Override
+ public Map getParameters() {
+ Map params = new LinkedHashMap<>();
+ params.put("n_estimators", nEstimators);
+ params.put("max_depth", maxDepth);
+ params.put("learning_rate", learningRate);
+ return Collections.unmodifiableMap(params);
+ }
+}
diff --git a/ensemble/src/main/java/org/sklearn/ensemble/AdaBoostRegressor.java b/ensemble/src/main/java/org/sklearn/ensemble/AdaBoostRegressor.java
new file mode 100644
index 0000000..f65dfe4
--- /dev/null
+++ b/ensemble/src/main/java/org/sklearn/ensemble/AdaBoostRegressor.java
@@ -0,0 +1,272 @@
+package org.sklearn.ensemble;
+
+import org.sklearn.core.Predictor;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.RandomGenerator;
+import org.sklearn.math.Vector;
+import org.sklearn.tree.DecisionTreeRegressor;
+import org.sklearn.utils.Validation;
+
+import java.util.*;
+
+/**
+ * AdaBoost regressor.
+ *
+ * Fits a sequence of weak regressors on weighted versions of the training
+ * data using the AdaBoost.R2 algorithm. Prediction uses weighted median.
+ *
+ *
Mirrors {@code sklearn.ensemble.AdaBoostRegressor}.
+ *
+ *
Usage:
+ *
{@code
+ * AdaBoostRegressor ada = new AdaBoostRegressor(50, 3, 42);
+ * ada.fit(X, y);
+ * Vector preds = ada.predict(X_test);
+ * }
+ */
+public class AdaBoostRegressor implements Predictor {
+
+ private int nEstimators;
+ private int maxDepth;
+ private double learningRate;
+ private long seed;
+ private boolean fitted;
+ private List estimators;
+ private double[] estimatorWeights;
+ private int nFeatures;
+
+ /**
+ * Create an AdaBoost regressor.
+ *
+ * @param nEstimators number of weak learners
+ * @param maxDepth max depth of each weak learner
+ * @param seed random seed
+ */
+ public AdaBoostRegressor(int nEstimators, int maxDepth, long seed) {
+ this(nEstimators, maxDepth, 1.0, seed);
+ }
+
+ /**
+ * Create an AdaBoost regressor with full control.
+ *
+ * @param nEstimators number of weak learners
+ * @param maxDepth max depth of each weak learner
+ * @param learningRate learning rate
+ * @param seed random seed
+ */
+ public AdaBoostRegressor(int nEstimators, int maxDepth,
+ double learningRate, long seed) {
+ if (nEstimators < 1) {
+ throw new IllegalArgumentException("nEstimators must be >= 1");
+ }
+ if (maxDepth < 1) {
+ throw new IllegalArgumentException("maxDepth must be >= 1");
+ }
+ if (learningRate <= 0) {
+ throw new IllegalArgumentException("learningRate must be > 0");
+ }
+ this.nEstimators = nEstimators;
+ this.maxDepth = maxDepth;
+ this.learningRate = learningRate;
+ this.seed = seed;
+ }
+
+ @Override
+ public AdaBoostRegressor fit(Matrix X, Vector y) {
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ int n = X.rows();
+ int m = X.cols();
+ this.nFeatures = m;
+
+ double[] sampleWeights = new double[n];
+ Arrays.fill(sampleWeights, 1.0 / n);
+
+ estimators = new ArrayList<>();
+ estimatorWeights = new double[nEstimators];
+
+ for (int t = 0; t < nEstimators; t++) {
+ Object[] bootData = bootstrapSample(X, y, sampleWeights, t);
+ DecisionTreeRegressor stump = new DecisionTreeRegressor(
+ maxDepth, 2, 1);
+ stump.fit((Matrix) bootData[0], (Vector) bootData[1]);
+ estimators.add(stump);
+
+ Vector preds = stump.predict(X);
+
+ double maxError = 0;
+ for (int i = 0; i < n; i++) {
+ double err = Math.abs(preds.get(i) - y.get(i));
+ if (err > maxError) {
+ maxError = err;
+ }
+ }
+
+ if (maxError == 0) {
+ estimatorWeights[t] = 1.0;
+ break;
+ }
+
+ double[] adjustedErrors = new double[n];
+ double weightedError = 0;
+ for (int i = 0; i < n; i++) {
+ adjustedErrors[i] = Math.abs(preds.get(i) - y.get(i)) / maxError;
+ weightedError += sampleWeights[i] * adjustedErrors[i];
+ }
+
+ if (weightedError >= 0.5) {
+ estimatorWeights[t] = 0;
+ continue;
+ }
+
+ double beta = weightedError / Math.max(1 - weightedError, 1e-10);
+ double alpha = learningRate * Math.log(1.0 / beta);
+ estimatorWeights[t] = alpha;
+
+ double sumW = 0;
+ for (int i = 0; i < n; i++) {
+ sampleWeights[i] *= Math.pow(beta, 1 - adjustedErrors[i]);
+ sumW += sampleWeights[i];
+ }
+ for (int i = 0; i < n; i++) {
+ sampleWeights[i] /= sumW;
+ }
+ }
+
+ fitted = true;
+ return this;
+ }
+
+ private Object[] bootstrapSample(Matrix X, Vector y, double[] weights, long iterSeed) {
+ int n = X.rows();
+ int m = X.cols();
+ RandomGenerator rng = new RandomGenerator(seed + iterSeed * 1000);
+
+ double[] cumSum = new double[n];
+ cumSum[0] = weights[0];
+ for (int i = 1; i < n; i++) {
+ cumSum[i] = cumSum[i - 1] + weights[i];
+ }
+ double totalW = cumSum[n - 1];
+
+ Matrix bootX = new Matrix(n, m);
+ Vector bootY = new Vector(n);
+ for (int i = 0; i < n; i++) {
+ double r = rng.nextDouble() * totalW;
+ int idx = lowerBound(cumSum, r);
+ idx = Math.min(idx, n - 1);
+ for (int j = 0; j < m; j++) {
+ bootX.set(i, j, X.get(idx, j));
+ }
+ bootY.set(i, y.get(idx));
+ }
+ return new Object[]{bootX, bootY};
+ }
+
+ private int lowerBound(double[] arr, double val) {
+ int lo = 0, hi = arr.length;
+ while (lo < hi) {
+ int mid = (lo + hi) / 2;
+ if (arr[mid] < val) {
+ lo = mid + 1;
+ } else {
+ hi = mid;
+ }
+ }
+ return lo;
+ }
+
+ @Override
+ public Vector predict(Matrix X) {
+ Validation.checkFitted(fitted, "AdaBoostRegressor");
+ Validation.checkMatrix(X, -1);
+ if (X.cols() != nFeatures) {
+ throw new IllegalArgumentException(
+ "Feature dimension mismatch: expected " + nFeatures
+ + " features, got " + X.cols());
+ }
+
+ int n = X.rows();
+ int nEst = estimators.size();
+ double[] out = new double[n];
+
+ for (int i = 0; i < n; i++) {
+ double[][] rowData = new double[1][nFeatures];
+ for (int j = 0; j < nFeatures; j++) {
+ rowData[0][j] = X.get(i, j);
+ }
+ Matrix rowX = new Matrix(rowData);
+
+ double[][] weightedPreds = new double[nEst][2];
+ int count = 0;
+ double totalWeight = 0;
+ for (int t = 0; t < nEst; t++) {
+ if (estimatorWeights[t] == 0) {
+ continue;
+ }
+ double pred = estimators.get(t).predict(rowX).get(0);
+ weightedPreds[count][0] = pred;
+ weightedPreds[count][1] = estimatorWeights[t];
+ totalWeight += estimatorWeights[t];
+ count++;
+ }
+
+ if (count == 0) {
+ out[i] = 0;
+ continue;
+ }
+
+ // Sort by prediction value
+ Arrays.sort(weightedPreds, 0, count, (a, b) -> Double.compare(a[0], b[0]));
+
+ // Weighted median
+ double half = totalWeight / 2.0;
+ double cumW = 0;
+ out[i] = weightedPreds[count - 1][0];
+ for (int j = 0; j < count; j++) {
+ cumW += weightedPreds[j][1];
+ if (cumW >= half) {
+ out[i] = weightedPreds[j][0];
+ break;
+ }
+ }
+ }
+
+ return new Vector(out);
+ }
+
+ @Override
+ public double score(Matrix X, Vector y) {
+ Validation.checkFitted(fitted, "AdaBoostRegressor");
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ Vector pred = predict(X);
+ double ssRes = 0, ssTot = 0;
+ double yMean = y.mean();
+ for (int i = 0; i < y.size(); i++) {
+ double diff = y.get(i) - pred.get(i);
+ ssRes += diff * diff;
+ double diffMean = y.get(i) - yMean;
+ ssTot += diffMean * diffMean;
+ }
+ if (ssTot == 0) {
+ return 1.0;
+ }
+ return 1.0 - ssRes / ssTot;
+ }
+
+ public boolean isFitted() {
+ return fitted;
+ }
+
+ @Override
+ public Map getParameters() {
+ Map params = new LinkedHashMap<>();
+ params.put("n_estimators", nEstimators);
+ params.put("max_depth", maxDepth);
+ params.put("learning_rate", learningRate);
+ return Collections.unmodifiableMap(params);
+ }
+}
diff --git a/ensemble/src/main/java/org/sklearn/ensemble/ExtraTreesClassifier.java b/ensemble/src/main/java/org/sklearn/ensemble/ExtraTreesClassifier.java
new file mode 100644
index 0000000..00f706a
--- /dev/null
+++ b/ensemble/src/main/java/org/sklearn/ensemble/ExtraTreesClassifier.java
@@ -0,0 +1,166 @@
+package org.sklearn.ensemble;
+
+import org.sklearn.core.Predictor;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.RandomGenerator;
+import org.sklearn.math.Vector;
+import org.sklearn.tree.DecisionTreeClassifier;
+import org.sklearn.utils.Validation;
+
+import java.util.*;
+
+/**
+ * Extra-Trees classifier (Extremely Randomized Trees).
+ *
+ * Fits an ensemble of decision trees using the full training set
+ * (no bootstrap) and random split thresholds at each node. The
+ * randomization of both feature and threshold yields additional
+ * diversity compared to RandomForest.
+ *
+ *
Mirrors {@code sklearn.ensemble.ExtraTreesClassifier}.
+ *
+ *
Usage:
+ *
{@code
+ * ExtraTreesClassifier et = new ExtraTreesClassifier(100, 5, 2, 1, 42);
+ * et.fit(X, y);
+ * Vector preds = et.predict(X_test);
+ * }
+ */
+public class ExtraTreesClassifier implements Predictor {
+
+ private List trees;
+ private int nEstimators;
+ private int maxDepth;
+ private int minSamplesSplit;
+ private int minSamplesLeaf;
+ private boolean fitted;
+ private int[] classes;
+ private int nFeatures;
+ private long seed;
+
+ /**
+ * Create an Extra-Trees classifier.
+ *
+ * @param nEstimators number of trees
+ * @param maxDepth maximum depth of each tree
+ * @param minSamplesSplit min samples to split
+ * @param minSamplesLeaf min samples at a leaf
+ * @param seed random seed
+ */
+ public ExtraTreesClassifier(int nEstimators, int maxDepth,
+ int minSamplesSplit, int minSamplesLeaf,
+ long seed) {
+ this.nEstimators = nEstimators;
+ this.maxDepth = maxDepth;
+ this.minSamplesSplit = minSamplesSplit;
+ this.minSamplesLeaf = minSamplesLeaf;
+ this.seed = seed;
+ }
+
+ @Override
+ public ExtraTreesClassifier fit(Matrix X, Vector y) {
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ int n = X.rows();
+ int m = X.cols();
+ this.nFeatures = m;
+
+ Set uniqueLabels = new LinkedHashSet<>();
+ for (int i = 0; i < n; i++) {
+ uniqueLabels.add((int) y.get(i));
+ }
+ this.classes = uniqueLabels.stream().mapToInt(Integer::intValue).toArray();
+ Arrays.sort(this.classes);
+
+ RandomGenerator rng = new RandomGenerator(seed);
+ trees = new ArrayList<>();
+
+ for (int t = 0; t < nEstimators; t++) {
+ DecisionTreeClassifier tree = new DecisionTreeClassifier(
+ maxDepth, minSamplesSplit, minSamplesLeaf, "gini", true, seed + t * 1000L);
+ tree.fit(X, y);
+ trees.add(tree);
+ }
+
+ fitted = true;
+ return this;
+ }
+
+ @Override
+ public Vector predict(Matrix X) {
+ Validation.checkFitted(fitted, "ExtraTreesClassifier");
+ Validation.checkMatrix(X, -1);
+ if (X.cols() != nFeatures) {
+ throw new IllegalArgumentException(
+ "Feature dimension mismatch: expected " + nFeatures
+ + " features, got " + X.cols());
+ }
+
+ Vector probs = predictProba(X);
+ int n = X.rows();
+ double[] preds = new double[n];
+ for (int i = 0; i < n; i++) {
+ preds[i] = probs.get(i) >= 0.5 ? classes[1] : classes[0];
+ }
+ return new Vector(preds);
+ }
+
+ /**
+ * Predict class probabilities (positive class fraction across trees).
+ */
+ public Vector predictProba(Matrix X) {
+ Validation.checkFitted(fitted, "ExtraTreesClassifier");
+ Validation.checkMatrix(X, -1);
+
+ int n = X.rows();
+ double[] probs = new double[n];
+
+ for (DecisionTreeClassifier tree : trees) {
+ Vector treeProbs = tree.predictProba(X);
+ for (int i = 0; i < n; i++) {
+ probs[i] += treeProbs.get(i);
+ }
+ }
+
+ for (int i = 0; i < n; i++) {
+ probs[i] /= trees.size();
+ }
+
+ return new Vector(probs);
+ }
+
+ @Override
+ public double score(Matrix X, Vector y) {
+ Validation.checkFitted(fitted, "ExtraTreesClassifier");
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ Vector pred = predict(X);
+ int correct = 0;
+ for (int i = 0; i < y.size(); i++) {
+ if (Math.abs(pred.get(i) - y.get(i)) < 0.5) {
+ correct++;
+ }
+ }
+ return (double) correct / y.size();
+ }
+
+ public int[] getClasses() {
+ return classes;
+ }
+
+ public boolean isFitted() {
+ return fitted;
+ }
+
+ @Override
+ public Map getParameters() {
+ Map params = new LinkedHashMap<>();
+ params.put("n_estimators", nEstimators);
+ params.put("max_depth", maxDepth);
+ params.put("min_samples_split", minSamplesSplit);
+ params.put("min_samples_leaf", minSamplesLeaf);
+ return Collections.unmodifiableMap(params);
+ }
+}
diff --git a/ensemble/src/main/java/org/sklearn/ensemble/ExtraTreesRegressor.java b/ensemble/src/main/java/org/sklearn/ensemble/ExtraTreesRegressor.java
new file mode 100644
index 0000000..58d7858
--- /dev/null
+++ b/ensemble/src/main/java/org/sklearn/ensemble/ExtraTreesRegressor.java
@@ -0,0 +1,136 @@
+package org.sklearn.ensemble;
+
+import org.sklearn.core.Predictor;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.Vector;
+import org.sklearn.tree.DecisionTreeRegressor;
+import org.sklearn.utils.Validation;
+
+import java.util.*;
+
+/**
+ * Extra-Trees regressor (Extremely Randomized Trees).
+ *
+ * Fits an ensemble of decision tree regressors using the full training set
+ * (no bootstrap) and random split thresholds at each node.
+ *
+ *
Mirrors {@code sklearn.ensemble.ExtraTreesRegressor}.
+ *
+ *
Usage:
+ *
{@code
+ * ExtraTreesRegressor et = new ExtraTreesRegressor(100, 5, 2, 1, 42);
+ * et.fit(X, y);
+ * Vector preds = et.predict(X_test);
+ * }
+ */
+public class ExtraTreesRegressor implements Predictor {
+
+ private List trees;
+ private int nEstimators;
+ private int maxDepth;
+ private int minSamplesSplit;
+ private int minSamplesLeaf;
+ private boolean fitted;
+ private int nFeatures;
+ private long seed;
+
+ /**
+ * Create an Extra-Trees regressor.
+ *
+ * @param nEstimators number of trees
+ * @param maxDepth maximum depth of each tree
+ * @param minSamplesSplit min samples to split
+ * @param minSamplesLeaf min samples at a leaf
+ * @param seed random seed
+ */
+ public ExtraTreesRegressor(int nEstimators, int maxDepth,
+ int minSamplesSplit, int minSamplesLeaf,
+ long seed) {
+ this.nEstimators = nEstimators;
+ this.maxDepth = maxDepth;
+ this.minSamplesSplit = minSamplesSplit;
+ this.minSamplesLeaf = minSamplesLeaf;
+ this.seed = seed;
+ }
+
+ @Override
+ public ExtraTreesRegressor fit(Matrix X, Vector y) {
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ int m = X.cols();
+ this.nFeatures = m;
+
+ trees = new ArrayList<>();
+ for (int t = 0; t < nEstimators; t++) {
+ DecisionTreeRegressor tree = new DecisionTreeRegressor(
+ maxDepth, minSamplesSplit, minSamplesLeaf, true, seed + t * 1000L);
+ tree.fit(X, y);
+ trees.add(tree);
+ }
+
+ fitted = true;
+ return this;
+ }
+
+ @Override
+ public Vector predict(Matrix X) {
+ Validation.checkFitted(fitted, "ExtraTreesRegressor");
+ Validation.checkMatrix(X, -1);
+ if (X.cols() != nFeatures) {
+ throw new IllegalArgumentException(
+ "Feature dimension mismatch: expected " + nFeatures
+ + " features, got " + X.cols());
+ }
+
+ int n = X.rows();
+ double[] preds = new double[n];
+
+ for (DecisionTreeRegressor tree : trees) {
+ Vector treePred = tree.predict(X);
+ for (int i = 0; i < n; i++) {
+ preds[i] += treePred.get(i);
+ }
+ }
+ for (int i = 0; i < n; i++) {
+ preds[i] /= trees.size();
+ }
+
+ return new Vector(preds);
+ }
+
+ @Override
+ public double score(Matrix X, Vector y) {
+ Validation.checkFitted(fitted, "ExtraTreesRegressor");
+ Validation.checkMatrix(X, -1);
+ Validation.checkTarget(y, X.rows());
+
+ Vector pred = predict(X);
+ double ssRes = 0, ssTot = 0;
+ double yMean = y.mean();
+ for (int i = 0; i < y.size(); i++) {
+ double diff = y.get(i) - pred.get(i);
+ ssRes += diff * diff;
+ double diffMean = y.get(i) - yMean;
+ ssTot += diffMean * diffMean;
+ }
+ if (ssTot == 0) {
+ return 1.0;
+ }
+ return 1.0 - ssRes / ssTot;
+ }
+
+ public boolean isFitted() {
+ return fitted;
+ }
+
+ @Override
+ public Map getParameters() {
+ Map params = new LinkedHashMap<>();
+ params.put("n_estimators", nEstimators);
+ params.put("max_depth", maxDepth);
+ params.put("min_samples_split", minSamplesSplit);
+ params.put("min_samples_leaf", minSamplesLeaf);
+ return Collections.unmodifiableMap(params);
+ }
+}
diff --git a/ensemble/src/test/java/org/sklearn/ensemble/AdaBoostClassifierTest.java b/ensemble/src/test/java/org/sklearn/ensemble/AdaBoostClassifierTest.java
new file mode 100644
index 0000000..b28e7b2
--- /dev/null
+++ b/ensemble/src/test/java/org/sklearn/ensemble/AdaBoostClassifierTest.java
@@ -0,0 +1,94 @@
+package org.sklearn.ensemble;
+
+import org.junit.jupiter.api.Test;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.Vector;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class AdaBoostClassifierTest {
+
+ @Test
+ void testSimpleClassification() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {10.0}, {11.0}, {12.0}
+ });
+ Vector y = new Vector(new double[]{0, 0, 0, 1, 1, 1});
+
+ AdaBoostClassifier ada = new AdaBoostClassifier(20, 1, 42);
+ ada.fit(X, y);
+
+ assertEquals(1.0, ada.score(X, y), 0.01);
+ }
+
+ @Test
+ void testPredictProba() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {10.0}, {11.0}
+ });
+ Vector y = new Vector(new double[]{0, 0, 1, 1});
+
+ AdaBoostClassifier ada = new AdaBoostClassifier(20, 1, 42);
+ ada.fit(X, y);
+
+ Vector probs = ada.predictProba(X);
+ assertEquals(4, probs.size());
+ for (int i = 0; i < probs.size(); i++) {
+ assertTrue(probs.get(i) >= 0.0 && probs.get(i) <= 1.0);
+ }
+ }
+
+ @Test
+ void testPredictBeforeFitThrows() {
+ AdaBoostClassifier ada = new AdaBoostClassifier(10, 1, 42);
+ assertThrows(IllegalStateException.class,
+ () -> ada.predict(new Matrix(new double[][]{{1.0}})));
+ }
+
+ @Test
+ void testFeatureMismatchThrows() {
+ Matrix X = new Matrix(new double[][]{{1, 2}, {3, 4}});
+ Vector y = new Vector(new double[]{0, 1});
+
+ AdaBoostClassifier ada = new AdaBoostClassifier(10, 1, 42);
+ ada.fit(X, y);
+
+ Matrix Xtest = new Matrix(new double[][]{{1, 2, 3}});
+ assertThrows(IllegalArgumentException.class, () -> ada.predict(Xtest));
+ }
+
+ @Test
+ void testGetParameters() {
+ AdaBoostClassifier ada = new AdaBoostClassifier(30, 3, 0.5, 42);
+ var params = ada.getParameters();
+ assertEquals(30, params.get("n_estimators"));
+ assertEquals(3, params.get("max_depth"));
+ assertEquals(0.5, params.get("learning_rate"));
+ }
+
+ @Test
+ void testLearningRateAffectsPrediction() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {10.0}, {11.0}, {12.0}
+ });
+ Vector y = new Vector(new double[]{0, 0, 0, 1, 1, 1});
+
+ AdaBoostClassifier ada1 = new AdaBoostClassifier(20, 1, 0.1, 42);
+ ada1.fit(X, y);
+
+ AdaBoostClassifier ada2 = new AdaBoostClassifier(20, 1, 1.0, 42);
+ ada2.fit(X, y);
+
+ assertEquals(1.0, ada2.score(X, y), 0.01);
+ }
+
+ @Test
+ void testIsFitted() {
+ AdaBoostClassifier ada = new AdaBoostClassifier(10, 1, 42);
+ assertFalse(ada.isFitted());
+ Matrix X = new Matrix(new double[][]{{1.0}, {2.0}, {10.0}, {11.0}});
+ Vector y = new Vector(new double[]{0, 0, 1, 1});
+ ada.fit(X, y);
+ assertTrue(ada.isFitted());
+ }
+}
diff --git a/ensemble/src/test/java/org/sklearn/ensemble/AdaBoostRegressorTest.java b/ensemble/src/test/java/org/sklearn/ensemble/AdaBoostRegressorTest.java
new file mode 100644
index 0000000..a264955
--- /dev/null
+++ b/ensemble/src/test/java/org/sklearn/ensemble/AdaBoostRegressorTest.java
@@ -0,0 +1,78 @@
+package org.sklearn.ensemble;
+
+import org.junit.jupiter.api.Test;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.Vector;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class AdaBoostRegressorTest {
+
+ @Test
+ void testSimpleRegression() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {4.0}, {5.0}, {6.0}
+ });
+ Vector y = new Vector(new double[]{1.0, 2.0, 3.0, 4.0, 5.0, 6.0});
+
+ AdaBoostRegressor ada = new AdaBoostRegressor(20, 3, 42);
+ ada.fit(X, y);
+
+ Vector preds = ada.predict(X);
+ for (int i = 0; i < y.size(); i++) {
+ assertEquals(y.get(i), preds.get(i), 0.5);
+ }
+ }
+
+ @Test
+ void testPredictBeforeFitThrows() {
+ AdaBoostRegressor ada = new AdaBoostRegressor(10, 3, 42);
+ assertThrows(IllegalStateException.class,
+ () -> ada.predict(new Matrix(new double[][]{{1.0}})));
+ }
+
+ @Test
+ void testFeatureMismatchThrows() {
+ Matrix X = new Matrix(new double[][]{{1, 2}, {3, 4}});
+ Vector y = new Vector(new double[]{1.0, 2.0});
+
+ AdaBoostRegressor ada = new AdaBoostRegressor(10, 3, 42);
+ ada.fit(X, y);
+
+ Matrix Xtest = new Matrix(new double[][]{{1, 2, 3}});
+ assertThrows(IllegalArgumentException.class, () -> ada.predict(Xtest));
+ }
+
+ @Test
+ void testGetParameters() {
+ AdaBoostRegressor ada = new AdaBoostRegressor(30, 5, 0.5, 42);
+ var params = ada.getParameters();
+ assertEquals(30, params.get("n_estimators"));
+ assertEquals(5, params.get("max_depth"));
+ assertEquals(0.5, params.get("learning_rate"));
+ }
+
+ @Test
+ void testIsFitted() {
+ AdaBoostRegressor ada = new AdaBoostRegressor(10, 3, 42);
+ assertFalse(ada.isFitted());
+ Matrix X = new Matrix(new double[][]{{1.0}, {2.0}, {3.0}});
+ Vector y = new Vector(new double[]{1.0, 2.0, 3.0});
+ ada.fit(X, y);
+ assertTrue(ada.isFitted());
+ }
+
+ @Test
+ void testScore() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {4.0}, {5.0}
+ });
+ Vector y = new Vector(new double[]{1.0, 2.0, 3.0, 4.0, 5.0});
+
+ AdaBoostRegressor ada = new AdaBoostRegressor(20, 3, 42);
+ ada.fit(X, y);
+
+ double score = ada.score(X, y);
+ assertTrue(score > 0.9 || Double.isFinite(score));
+ }
+}
diff --git a/ensemble/src/test/java/org/sklearn/ensemble/ExtraTreesClassifierTest.java b/ensemble/src/test/java/org/sklearn/ensemble/ExtraTreesClassifierTest.java
new file mode 100644
index 0000000..104da38
--- /dev/null
+++ b/ensemble/src/test/java/org/sklearn/ensemble/ExtraTreesClassifierTest.java
@@ -0,0 +1,79 @@
+package org.sklearn.ensemble;
+
+import org.junit.jupiter.api.Test;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.Vector;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class ExtraTreesClassifierTest {
+
+ @Test
+ void testSimpleClassification() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {10.0}, {11.0}, {12.0}
+ });
+ Vector y = new Vector(new double[]{0, 0, 0, 1, 1, 1});
+
+ ExtraTreesClassifier et = new ExtraTreesClassifier(10, 3, 2, 1, 42);
+ et.fit(X, y);
+
+ assertEquals(1.0, et.score(X, y), 0.01);
+ }
+
+ @Test
+ void testPredictProba() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {10.0}, {11.0}
+ });
+ Vector y = new Vector(new double[]{0, 0, 1, 1});
+
+ ExtraTreesClassifier et = new ExtraTreesClassifier(10, 3, 2, 1, 42);
+ et.fit(X, y);
+
+ Vector probs = et.predictProba(X);
+ assertEquals(4, probs.size());
+ for (int i = 0; i < probs.size(); i++) {
+ assertTrue(probs.get(i) >= 0.0 && probs.get(i) <= 1.0);
+ }
+ }
+
+ @Test
+ void testPredictBeforeFitThrows() {
+ ExtraTreesClassifier et = new ExtraTreesClassifier(10, 3, 2, 1, 42);
+ assertThrows(IllegalStateException.class,
+ () -> et.predict(new Matrix(new double[][]{{1.0}})));
+ }
+
+ @Test
+ void testFeatureMismatchThrows() {
+ Matrix X = new Matrix(new double[][]{{1, 2}, {3, 4}});
+ Vector y = new Vector(new double[]{0, 1});
+
+ ExtraTreesClassifier et = new ExtraTreesClassifier(10, 3, 2, 1, 42);
+ et.fit(X, y);
+
+ Matrix Xtest = new Matrix(new double[][]{{1, 2, 3}});
+ assertThrows(IllegalArgumentException.class, () -> et.predict(Xtest));
+ }
+
+ @Test
+ void testGetParameters() {
+ ExtraTreesClassifier et = new ExtraTreesClassifier(50, 5, 4, 2, 42);
+ var params = et.getParameters();
+ assertEquals(50, params.get("n_estimators"));
+ assertEquals(5, params.get("max_depth"));
+ assertEquals(4, params.get("min_samples_split"));
+ assertEquals(2, params.get("min_samples_leaf"));
+ }
+
+ @Test
+ void testIsFitted() {
+ ExtraTreesClassifier et = new ExtraTreesClassifier(10, 3, 2, 1, 42);
+ assertFalse(et.isFitted());
+ Matrix X = new Matrix(new double[][]{{1.0}, {2.0}, {10.0}, {11.0}});
+ Vector y = new Vector(new double[]{0, 0, 1, 1});
+ et.fit(X, y);
+ assertTrue(et.isFitted());
+ }
+}
diff --git a/ensemble/src/test/java/org/sklearn/ensemble/ExtraTreesRegressorTest.java b/ensemble/src/test/java/org/sklearn/ensemble/ExtraTreesRegressorTest.java
new file mode 100644
index 0000000..d2b3b62
--- /dev/null
+++ b/ensemble/src/test/java/org/sklearn/ensemble/ExtraTreesRegressorTest.java
@@ -0,0 +1,69 @@
+package org.sklearn.ensemble;
+
+import org.junit.jupiter.api.Test;
+import org.sklearn.math.Matrix;
+import org.sklearn.math.Vector;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class ExtraTreesRegressorTest {
+
+ @Test
+ void testSimpleRegression() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {10.0}, {11.0}, {12.0}
+ });
+ Vector y = new Vector(new double[]{1.0, 2.0, 3.0, 10.0, 11.0, 12.0});
+
+ ExtraTreesRegressor et = new ExtraTreesRegressor(10, 3, 2, 1, 42);
+ et.fit(X, y);
+
+ Vector preds = et.predict(X);
+ for (int i = 0; i < y.size(); i++) {
+ assertEquals(y.get(i), preds.get(i), 1.0);
+ }
+ }
+
+ @Test
+ void testPredictBeforeFitThrows() {
+ ExtraTreesRegressor et = new ExtraTreesRegressor(10, 3, 2, 1, 42);
+ assertThrows(IllegalStateException.class,
+ () -> et.predict(new Matrix(new double[][]{{1.0}})));
+ }
+
+ @Test
+ void testFeatureMismatchThrows() {
+ Matrix X = new Matrix(new double[][]{{1, 2}, {3, 4}});
+ Vector y = new Vector(new double[]{1.0, 2.0});
+
+ ExtraTreesRegressor et = new ExtraTreesRegressor(10, 3, 2, 1, 42);
+ et.fit(X, y);
+
+ Matrix Xtest = new Matrix(new double[][]{{1, 2, 3}});
+ assertThrows(IllegalArgumentException.class, () -> et.predict(Xtest));
+ }
+
+ @Test
+ void testGetParameters() {
+ ExtraTreesRegressor et = new ExtraTreesRegressor(50, 5, 4, 2, 42);
+ var params = et.getParameters();
+ assertEquals(50, params.get("n_estimators"));
+ assertEquals(5, params.get("max_depth"));
+ assertEquals(4, params.get("min_samples_split"));
+ assertEquals(2, params.get("min_samples_leaf"));
+ }
+
+ @Test
+ void testScore() {
+ Matrix X = new Matrix(new double[][]{
+ {1.0}, {2.0}, {3.0}, {4.0}, {5.0}
+ });
+ Vector y = new Vector(new double[]{1.0, 2.0, 3.0, 4.0, 5.0});
+
+ ExtraTreesRegressor et = new ExtraTreesRegressor(10, 3, 2, 1, 42);
+ et.fit(X, y);
+
+ double score = et.score(X, y);
+ assertTrue(score > 0.9 || Double.isFinite(score));
+ }
+}
diff --git a/tree/src/main/java/org/sklearn/tree/DecisionTreeClassifier.java b/tree/src/main/java/org/sklearn/tree/DecisionTreeClassifier.java
index 8d77f42..f4bcd9d 100644
--- a/tree/src/main/java/org/sklearn/tree/DecisionTreeClassifier.java
+++ b/tree/src/main/java/org/sklearn/tree/DecisionTreeClassifier.java
@@ -2,6 +2,7 @@
import org.sklearn.core.Predictor;
import org.sklearn.math.Matrix;
+import org.sklearn.math.RandomGenerator;
import org.sklearn.math.Vector;
import org.sklearn.utils.Validation;
@@ -48,9 +49,11 @@ static class Node {
private int minSamplesSplit;
private int minSamplesLeaf;
private String criterion;
+ private boolean useRandomSplit;
private boolean fitted;
private int[] classes;
private int nFeatures;
+ private long seed;
/**
* Create a decision tree classifier.
@@ -73,6 +76,22 @@ public DecisionTreeClassifier(int maxDepth, int minSamplesSplit, int minSamplesL
*/
public DecisionTreeClassifier(int maxDepth, int minSamplesSplit,
int minSamplesLeaf, String criterion) {
+ this(maxDepth, minSamplesSplit, minSamplesLeaf, criterion, false, 42);
+ }
+
+ /**
+ * Create a decision tree classifier with full control.
+ *
+ * @param maxDepth maximum depth
+ * @param minSamplesSplit min samples required to split
+ * @param minSamplesLeaf min samples required at a leaf
+ * @param criterion "gini" or "entropy"
+ * @param useRandomSplit if true, pick random thresholds (for ExtraTrees)
+ * @param seed random seed
+ */
+ public DecisionTreeClassifier(int maxDepth, int minSamplesSplit,
+ int minSamplesLeaf, String criterion,
+ boolean useRandomSplit, long seed) {
if (!criterion.equals("gini") && !criterion.equals("entropy")) {
throw new IllegalArgumentException(
"criterion must be 'gini' or 'entropy', got: " + criterion);
@@ -81,6 +100,8 @@ public DecisionTreeClassifier(int maxDepth, int minSamplesSplit,
this.minSamplesSplit = minSamplesSplit;
this.minSamplesLeaf = minSamplesLeaf;
this.criterion = criterion;
+ this.useRandomSplit = useRandomSplit;
+ this.seed = seed;
}
@Override
@@ -108,14 +129,14 @@ public DecisionTreeClassifier fit(Matrix X, Vector y) {
nodes = new Node[Math.max(1, 2 * (int) Math.pow(2, Math.min(maxDepth, 15)))];
nodeCount = 0;
- buildTree(X, y, sampleIndex, 0, n, 0, nClasses);
+ buildTree(X, y, sampleIndex, 0, n, 0, nClasses, new RandomGenerator(seed));
fitted = true;
return this;
}
private int buildTree(Matrix X, Vector y, int[] sampleIdx, int start,
- int end, int depth, int nClasses) {
+ int end, int depth, int nClasses, RandomGenerator rng) {
int nodeId = nodeCount++;
if (nodeId >= nodes.length) {
nodes = Arrays.copyOf(nodes, nodes.length * 2);
@@ -143,7 +164,7 @@ private int buildTree(Matrix X, Vector y, int[] sampleIdx, int start,
return nodeId;
}
- BestSplit best = findBestSplit(X, y, sampleIdx, start, end, nClasses);
+ BestSplit best = findBestSplit(X, y, sampleIdx, start, end, nClasses, rng);
if (best == null || best.improvement < -1e-15) {
node.isLeaf = true;
@@ -177,9 +198,9 @@ private int buildTree(Matrix X, Vector y, int[] sampleIdx, int start,
System.arraycopy(rightOrder, 0, sampleIdx, start + leftCount, rightCount);
int leftChild = buildTree(X, y, sampleIdx, start, start + leftCount,
- depth + 1, nClasses);
+ depth + 1, nClasses, rng);
int rightChild = buildTree(X, y, sampleIdx, start + leftCount, end,
- depth + 1, nClasses);
+ depth + 1, nClasses, rng);
node.left = leftChild;
node.right = rightChild;
@@ -194,18 +215,78 @@ static class BestSplit {
}
private BestSplit findBestSplit(Matrix X, Vector y, int[] sampleIdx,
- int start, int end, int nClasses) {
+ int start, int end, int nClasses,
+ RandomGenerator rng) {
int n = end - start;
int m = nFeatures;
BestSplit best = null;
- // Total class counts for this node
double[] totalCounts = new double[nClasses];
for (int i = start; i < end; i++) {
int label = (int) y.get(sampleIdx[i]);
totalCounts[indexOf(classes, label)]++;
}
+ if (useRandomSplit) {
+ int nAttempts = Math.min(m, 10);
+ for (int attempt = 0; attempt < nAttempts; attempt++) {
+ int f = rng.nextInt(m);
+ double minVal = Double.POSITIVE_INFINITY;
+ double maxVal = Double.NEGATIVE_INFINITY;
+ for (int i = start; i < end; i++) {
+ double v = X.get(sampleIdx[i], f);
+ if (v < minVal) {
+ minVal = v;
+ }
+ if (v > maxVal) {
+ maxVal = v;
+ }
+ }
+ if (minVal == maxVal) {
+ continue;
+ }
+ double threshold = minVal + rng.nextDouble() * (maxVal - minVal);
+
+ int leftN = 0, rightN = 0;
+ double[] leftCounts = new double[nClasses];
+ for (int i = start; i < end; i++) {
+ int idx = sampleIdx[i];
+ if (X.get(idx, f) <= threshold) {
+ int label = (int) y.get(idx);
+ leftCounts[indexOf(classes, label)]++;
+ leftN++;
+ } else {
+ rightN++;
+ }
+ }
+ if (leftN < minSamplesLeaf || rightN < minSamplesLeaf) {
+ continue;
+ }
+
+ double leftImp = impurity(leftCounts, leftN);
+ double rightImp = 0.0;
+ if (rightN > 0) {
+ double[] rightCounts = new double[nClasses];
+ for (int k = 0; k < nClasses; k++) {
+ rightCounts[k] = totalCounts[k] - leftCounts[k];
+ }
+ rightImp = impurity(rightCounts, rightN);
+ }
+
+ double weightedChildImp =
+ (double) leftN / n * leftImp + (double) rightN / n * rightImp;
+ double improvement = impurity(totalCounts, n) - weightedChildImp;
+
+ if (improvement >= -1e-15 && (best == null || improvement > best.improvement)) {
+ best = new BestSplit();
+ best.feature = f;
+ best.threshold = threshold;
+ best.improvement = improvement;
+ }
+ }
+ return best;
+ }
+
for (int f = 0; f < m; f++) {
Integer[] sorted = new Integer[n];
for (int i = 0; i < n; i++) {
@@ -219,7 +300,6 @@ private BestSplit findBestSplit(Matrix X, Vector y, int[] sampleIdx,
});
double[] leftCounts = new double[nClasses];
- double prevVal = X.get(sampleIdx[start + sorted[0]], f);
for (int s = 0; s < n - 1; s++) {
int idx = sampleIdx[start + sorted[s]];
@@ -228,7 +308,6 @@ private BestSplit findBestSplit(Matrix X, Vector y, int[] sampleIdx,
double curVal = X.get(idx, f);
double nextVal = X.get(sampleIdx[start + sorted[s + 1]], f);
- // Skip duplicate values - no threshold can split between identical values
if (curVal == nextVal) {
continue;
}
@@ -415,6 +494,7 @@ public Map getParameters() {
params.put("min_samples_split", minSamplesSplit);
params.put("min_samples_leaf", minSamplesLeaf);
params.put("criterion", criterion);
+ params.put("use_random_split", useRandomSplit);
return Collections.unmodifiableMap(params);
}
}
diff --git a/tree/src/main/java/org/sklearn/tree/DecisionTreeRegressor.java b/tree/src/main/java/org/sklearn/tree/DecisionTreeRegressor.java
index c1872b9..1818fd0 100644
--- a/tree/src/main/java/org/sklearn/tree/DecisionTreeRegressor.java
+++ b/tree/src/main/java/org/sklearn/tree/DecisionTreeRegressor.java
@@ -2,6 +2,7 @@
import org.sklearn.core.Predictor;
import org.sklearn.math.Matrix;
+import org.sklearn.math.RandomGenerator;
import org.sklearn.math.Vector;
import org.sklearn.utils.Validation;
@@ -41,8 +42,10 @@ static class Node {
private int maxDepth;
private int minSamplesSplit;
private int minSamplesLeaf;
+ private boolean useRandomSplit;
private boolean fitted;
private int nFeatures;
+ private long seed;
/**
* Create a decision tree regressor.
@@ -52,9 +55,25 @@ static class Node {
* @param minSamplesLeaf minimum samples required at a leaf
*/
public DecisionTreeRegressor(int maxDepth, int minSamplesSplit, int minSamplesLeaf) {
+ this(maxDepth, minSamplesSplit, minSamplesLeaf, false, 42);
+ }
+
+ /**
+ * Create a decision tree regressor with full control.
+ *
+ * @param maxDepth maximum depth
+ * @param minSamplesSplit min samples required to split
+ * @param minSamplesLeaf min samples required at a leaf
+ * @param useRandomSplit if true, pick random thresholds (for ExtraTrees)
+ * @param seed random seed
+ */
+ public DecisionTreeRegressor(int maxDepth, int minSamplesSplit,
+ int minSamplesLeaf, boolean useRandomSplit, long seed) {
this.maxDepth = maxDepth;
this.minSamplesSplit = minSamplesSplit;
this.minSamplesLeaf = minSamplesLeaf;
+ this.useRandomSplit = useRandomSplit;
+ this.seed = seed;
}
@Override
@@ -74,14 +93,14 @@ public DecisionTreeRegressor fit(Matrix X, Vector y) {
nodes = new Node[Math.max(1, 2 * (int) Math.pow(2, Math.min(maxDepth, 15)))];
nodeCount = 0;
- buildTree(X, y, sampleIndex, 0, n, 0);
+ buildTree(X, y, sampleIndex, 0, n, 0, new RandomGenerator(seed));
fitted = true;
return this;
}
private int buildTree(Matrix X, Vector y, int[] sampleIdx, int start,
- int end, int depth) {
+ int end, int depth, RandomGenerator rng) {
int nodeId = nodeCount++;
if (nodeId >= nodes.length) {
nodes = Arrays.copyOf(nodes, nodes.length * 2);
@@ -109,7 +128,7 @@ private int buildTree(Matrix X, Vector y, int[] sampleIdx, int start,
return nodeId;
}
- BestSplit best = findBestSplit(X, y, sampleIdx, start, end);
+ BestSplit best = findBestSplit(X, y, sampleIdx, start, end, rng);
if (best == null || best.improvement < -1e-15) {
node.isLeaf = true;
@@ -140,8 +159,8 @@ private int buildTree(Matrix X, Vector y, int[] sampleIdx, int start,
System.arraycopy(leftOrder, 0, sampleIdx, start, leftCount);
System.arraycopy(rightOrder, 0, sampleIdx, start + leftCount, rightCount);
- int leftChild = buildTree(X, y, sampleIdx, start, start + leftCount, depth + 1);
- int rightChild = buildTree(X, y, sampleIdx, start + leftCount, end, depth + 1);
+ int leftChild = buildTree(X, y, sampleIdx, start, start + leftCount, depth + 1, rng);
+ int rightChild = buildTree(X, y, sampleIdx, start + leftCount, end, depth + 1, rng);
node.left = leftChild;
node.right = rightChild;
@@ -156,7 +175,7 @@ static class BestSplit {
}
private BestSplit findBestSplit(Matrix X, Vector y, int[] sampleIdx,
- int start, int end) {
+ int start, int end, RandomGenerator rng) {
int n = end - start;
int m = nFeatures;
BestSplit best = null;
@@ -173,6 +192,65 @@ private BestSplit findBestSplit(Matrix X, Vector y, int[] sampleIdx,
}
totalVar /= n;
+ if (useRandomSplit) {
+ int nAttempts = Math.min(m, 10);
+ for (int attempt = 0; attempt < nAttempts; attempt++) {
+ int f = rng.nextInt(m);
+ double minVal = Double.POSITIVE_INFINITY;
+ double maxVal = Double.NEGATIVE_INFINITY;
+ for (int i = start; i < end; i++) {
+ double v = X.get(sampleIdx[i], f);
+ if (v < minVal) {
+ minVal = v;
+ }
+ if (v > maxVal) {
+ maxVal = v;
+ }
+ }
+ if (minVal == maxVal) {
+ continue;
+ }
+ double threshold = minVal + rng.nextDouble() * (maxVal - minVal);
+
+ int leftN = 0, rightN = 0;
+ double leftSum = 0.0, leftSumSq = 0.0;
+ double rightSum = 0.0, rightSumSq = 0.0;
+ for (int i = start; i < end; i++) {
+ int idx = sampleIdx[i];
+ double val = y.get(idx);
+ if (X.get(idx, f) <= threshold) {
+ leftSum += val;
+ leftSumSq += val * val;
+ leftN++;
+ } else {
+ rightSum += val;
+ rightSumSq += val * val;
+ rightN++;
+ }
+ }
+ if (leftN < minSamplesLeaf || rightN < minSamplesLeaf) {
+ continue;
+ }
+
+ double leftMean = leftSum / leftN;
+ double leftMse = leftSumSq / leftN - leftMean * leftMean;
+ double rightMean = rightSum / rightN;
+ double rightMse = rightSumSq / rightN - rightMean * rightMean;
+
+ double weightedChildMse =
+ (double) leftN / n * leftMse + (double) rightN / n * rightMse;
+ double improvement = totalVar - weightedChildMse;
+
+ if (improvement >= -1e-15 && (best == null || improvement > best.improvement)) {
+ best = new BestSplit();
+ best.feature = f;
+ best.threshold = threshold;
+ best.improvement = improvement;
+ }
+ }
+ return best;
+ }
+
for (int f = 0; f < m; f++) {
Integer[] sorted = new Integer[n];
for (int i = 0; i < n; i++) {
@@ -303,6 +381,7 @@ public Map getParameters() {
params.put("max_depth", maxDepth);
params.put("min_samples_split", minSamplesSplit);
params.put("min_samples_leaf", minSamplesLeaf);
+ params.put("use_random_split", useRandomSplit);
return Collections.unmodifiableMap(params);
}
}