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); } }