Skip to content

Commit 9cdd051

Browse files
committed
Fixed related errors
1 parent 3fc82a7 commit 9cdd051

1 file changed

Lines changed: 44 additions & 16 deletions

File tree

‎src/main/java/com/thealgorithms/machinelearning/Perceptron.java‎

Lines changed: 44 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,9 @@
1313
* greater than or equal to zero, and {@code 0} otherwise. For a
1414
* misclassified sample, the update is {@code weight += learningRate * error *
1515
* feature} and {@code bias += learningRate * error}, where {@code error} is
16-
* the true label minus the prediction.
16+
* the true label minus the prediction. Samples are visited one at a time, so
17+
* each prediction uses the parameters produced by the preceding updates of
18+
* the same epoch.
1719
*
1820
* @see <a href="https://en.wikipedia.org/wiki/Perceptron">Perceptron</a>
1921
*/
@@ -22,7 +24,6 @@ public final class Perceptron {
2224
private final int maxEpochs;
2325
private double[] weights;
2426
private double bias;
25-
private int numFeatures;
2627
private int epochsRun;
2728
private boolean converged;
2829

@@ -48,18 +49,23 @@ public Perceptron(double learningRate, int maxEpochs) {
4849
* Fits the classifier using binary training labels.
4950
*
5051
* <p>Fitting resets the weights and bias to zero before training. The
51-
* method records whether an entire epoch completed without an update.
52+
* method records whether an entire epoch completed without an update. A
53+
* large {@code learningRate} combined with large feature values can push
54+
* the parameters past the range of {@code double}; the classifier then
55+
* returns to its unfitted state instead of reporting predictions derived
56+
* from non-finite parameters.
5257
*
5358
* @param features training feature vectors
5459
* @param labels corresponding binary labels, each either {@code 0} or
5560
* {@code 1}
5661
* @throws IllegalArgumentException if the training data is invalid
62+
* @throws ArithmeticException if training diverges and the learned
63+
* parameters stop being finite
5764
*/
5865
public void fit(double[][] features, int[] labels) {
5966
validateTrainingData(features, labels);
6067

61-
numFeatures = features[0].length;
62-
weights = new double[numFeatures];
68+
weights = new double[features[0].length];
6369
bias = 0.0;
6470
epochsRun = 0;
6571
converged = false;
@@ -68,7 +74,7 @@ public void fit(double[][] features, int[] labels) {
6874
boolean updated = false;
6975

7076
for (int sampleIndex = 0; sampleIndex < features.length; sampleIndex++) {
71-
int prediction = predict(features[sampleIndex]);
77+
int prediction = rawPredict(features[sampleIndex]);
7278
int error = labels[sampleIndex] - prediction;
7379

7480
if (error != 0) {
@@ -83,6 +89,8 @@ public void fit(double[][] features, int[] labels) {
8389
break;
8490
}
8591
}
92+
93+
ensureParametersAreFinite();
8694
}
8795

8896
/**
@@ -96,12 +104,7 @@ public void fit(double[][] features, int[] labels) {
96104
public int predict(double[] sample) {
97105
ensureFitted();
98106
validateSample(sample);
99-
100-
double weightedSum = bias;
101-
for (int featureIndex = 0; featureIndex < numFeatures; featureIndex++) {
102-
weightedSum += weights[featureIndex] * sample[featureIndex];
103-
}
104-
return weightedSum >= 0.0 ? 1 : 0;
107+
return rawPredict(sample);
105108
}
106109

107110
/**
@@ -170,13 +173,28 @@ public int getEpochsRun() {
170173
return epochsRun;
171174
}
172175

176+
private int rawPredict(double[] sample) {
177+
double weightedSum = bias;
178+
for (int featureIndex = 0; featureIndex < weights.length; featureIndex++) {
179+
weightedSum += weights[featureIndex] * sample[featureIndex];
180+
}
181+
return weightedSum >= 0.0 ? 1 : 0;
182+
}
183+
173184
private void update(double[] sample, int error) {
174-
for (int featureIndex = 0; featureIndex < numFeatures; featureIndex++) {
185+
for (int featureIndex = 0; featureIndex < weights.length; featureIndex++) {
175186
weights[featureIndex] += learningRate * error * sample[featureIndex];
176187
}
177188
bias += learningRate * error;
178189
}
179190

191+
private void ensureParametersAreFinite() {
192+
if (!Double.isFinite(bias) || !isFinite(weights)) {
193+
weights = null;
194+
throw new ArithmeticException("training diverged; try a smaller learningRate or scaled features");
195+
}
196+
}
197+
180198
private void ensureFitted() {
181199
if (weights == null) {
182200
throw new IllegalStateException("classifier has not been fitted");
@@ -200,7 +218,10 @@ private void validateTrainingData(double[][] features, int[] labels) {
200218
int featureCount = features[0].length;
201219
for (int sampleIndex = 0; sampleIndex < features.length; sampleIndex++) {
202220
double[] sample = features[sampleIndex];
203-
if (sample == null || sample.length != featureCount) {
221+
if (sample == null) {
222+
throw new IllegalArgumentException("feature vectors cannot be null or empty");
223+
}
224+
if (sample.length != featureCount) {
204225
throw new IllegalArgumentException("all feature vectors must have the same dimension");
205226
}
206227
validateFiniteValues(sample);
@@ -211,17 +232,24 @@ private void validateTrainingData(double[][] features, int[] labels) {
211232
}
212233

213234
private void validateSample(double[] sample) {
214-
if (sample == null || sample.length != numFeatures) {
235+
if (sample == null || sample.length != weights.length) {
215236
throw new IllegalArgumentException("sample must match the training feature dimension");
216237
}
217238
validateFiniteValues(sample);
218239
}
219240

220241
private static void validateFiniteValues(double[] values) {
242+
if (!isFinite(values)) {
243+
throw new IllegalArgumentException("feature values must be finite");
244+
}
245+
}
246+
247+
private static boolean isFinite(double[] values) {
221248
for (double value : values) {
222249
if (!Double.isFinite(value)) {
223-
throw new IllegalArgumentException("feature values must be finite");
250+
return false;
224251
}
225252
}
253+
return true;
226254
}
227255
}

0 commit comments

Comments
 (0)