Repository navigation
Expand file tree
/
Copy pathforest.ts
More file actions
157 lines (149 loc) · 4.55 KB
/
Copy pathforest.ts
File metadata and controls
157 lines (149 loc) · 4.55 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
import { concat, countBy, find, head, isEqual, keys, map, maxBy, range, reduce, values } from 'lodash';
import { DecisionTreeClassifier } from '../tree';
import { IMlModel, Type1DMatrix, Type2DMatrix } from '../types';
import { validateFitInputs, validateMatrix2D } from '../utils/validation';
/**
* Base RandomForest implementation used by both classifier and regressor
* @ignore
*/
export class BaseRandomForest implements IMlModel<number> {
protected trees = [];
protected nEstimator;
protected randomState = null;
/**
*
* @param {number} nEstimator - Number of trees.
* @param random_state - Random seed value for DecisionTrees
*/
constructor(
{
// Each object param default value
nEstimator = 10,
random_state = null,
}: {
// Param types
nEstimator?: number;
random_state?: number;
} = {
// Default value on empty constructor
nEstimator: 10,
random_state: null,
},
) {
this.nEstimator = nEstimator;
this.randomState = random_state;
}
/**
* Build a forest of trees from the training set (X, y).
* @param {Array} X - array-like or sparse matrix of shape = [n_samples, n_features]
* @param {Array} y - array-like, shape = [n_samples] or [n_samples, n_outputs]
* @returns void
*/
public fit(X: Type2DMatrix<number> = null, y: Type1DMatrix<number> = null): void {
validateFitInputs(X, y);
this.trees = reduce(
range(0, this.nEstimator),
(sum) => {
const tree = new DecisionTreeClassifier({
featureLabels: null,
random_state: this.randomState,
});
tree.fit(X, y);
return concat(sum, [tree]);
},
[],
);
}
/**
* Returning the current model's checkpoint
* @returns {{trees: any[]}}
*/
public toJSON(): {
/**
* Decision trees
*/
trees: any[];
} {
return {
trees: this.trees,
};
}
/**
* Restore the model from a checkpoint
* @param {any[]} trees - Decision trees
*/
public fromJSON({ trees = null }: { trees: any[] }): void {
if (!trees) {
throw new Error('You must provide both tree to restore the model');
}
this.trees = trees;
}
/**
* Internal predict function used by either RandomForestClassifier or Regressor
* @param X
* @private
*/
public predict(X: Type2DMatrix<number> = null): number[][] {
validateMatrix2D(X);
return map(this.trees, (tree: DecisionTreeClassifier) => {
// TODO: Check if it's a matrix or an array
return tree.predict(X);
});
}
}
/**
* Random forest classifier creates a set of decision trees from randomly selected subset of training set.
* It then aggregates the votes from different decision trees to decide the final class of the test object.
*
* @example
* import { RandomForestClassifier } from 'machinelearn/ensemble';
*
* const X = [[0, 0], [1, 1], [2, 1], [1, 5], [3, 2]];
* const y = [0, 1, 2, 3, 7];
*
* const randomForest = new RandomForestClassifier();
* randomForest.fit(X, y);
*
* // Results in a value such as [ '0', '2' ].
* // Predictions will change as we have not set a seed value.
*/
export class RandomForestClassifier extends BaseRandomForest {
/**
* Predict class for X.
*
* The predicted class of an input sample is a vote by the trees in the forest, weighted by their probability estimates.
* That is, the predicted class is the one with highest mean probability estimate across the trees.
* @param {Array} X - array-like or sparse matrix of shape = [n_samples]
* @returns {string[]}
*/
public predict(X: Type2DMatrix<number> = null): any[] {
const predictions = super.predict(X);
return this.votePredictions(predictions);
}
/**
* @hidden
* Bagging prediction helper method
* According to the predictions returns by the trees, it will select the
* class with the maximum number (votes)
* @param {Array<any>} predictions - List of initial predictions that may look like [ [1, 2], [1, 1] ... ]
* @returns {string[]}
*/
private votePredictions(predictions: Type2DMatrix<number>): number[] {
const counts = countBy(predictions, (x) => x);
const countsArray = reduce(
keys(counts),
(sum, k) => {
const returning = {};
returning[k] = counts[k];
return concat(sum, returning);
},
[],
);
const max = maxBy(countsArray, (x) => head(values(x)));
const key = head(keys(max));
// Find the actual class values from the predictions
return find(predictions, (pred) => {
return isEqual(pred.join(','), key);
});
}
}