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