Skip to content

Commit 3027ce1

Browse files
Th-Shivamalxkm
andauthored
Add Perceptron binary classifier (#7601)
* Add Perceptron binary classifier * Fixed related errors * Fix PMD static import violation and cover Perceptron gaps Reduce static imports to PMD's limit, assert the learned bias and weights, and cover the length-mismatch, null-vector, and divergence paths. --------- Co-authored-by: Alex <19151554+alxkm@users.noreply.github.com>
1 parent 4b58d36 commit 3027ce1

2 files changed

Lines changed: 418 additions & 0 deletions

File tree

Lines changed: 255 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,255 @@
1+
package com.thealgorithms.machinelearning;
2+
3+
/**
4+
* A binary Perceptron classifier.
5+
*
6+
* <p>The Perceptron is a single-layer neural network that learns a linear
7+
* decision boundary. It updates its weights whenever a training sample is
8+
* misclassified. Convergence is guaranteed for linearly separable data, but
9+
* training stops after the configured epoch limit for non-separable data.
10+
* Labels must be either {@code 0} or {@code 1}.
11+
*
12+
* <p>The prediction rule is {@code 1} when the weighted sum plus bias is
13+
* greater than or equal to zero, and {@code 0} otherwise. For a
14+
* misclassified sample, the update is {@code weight += learningRate * error *
15+
* feature} and {@code bias += learningRate * error}, where {@code error} is
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.
19+
*
20+
* @see <a href="https://en.wikipedia.org/wiki/Perceptron">Perceptron</a>
21+
*/
22+
public final class Perceptron {
23+
private final double learningRate;
24+
private final int maxEpochs;
25+
private double[] weights;
26+
private double bias;
27+
private int epochsRun;
28+
private boolean converged;
29+
30+
/**
31+
* Constructs a Perceptron with the given training hyperparameters.
32+
*
33+
* @param learningRate positive step size used for each update
34+
* @param maxEpochs positive maximum number of passes over the training data
35+
* @throws IllegalArgumentException if a hyperparameter is invalid
36+
*/
37+
public Perceptron(double learningRate, int maxEpochs) {
38+
if (!Double.isFinite(learningRate) || learningRate <= 0.0) {
39+
throw new IllegalArgumentException("learningRate must be finite and greater than 0");
40+
}
41+
if (maxEpochs <= 0) {
42+
throw new IllegalArgumentException("maxEpochs must be greater than 0");
43+
}
44+
this.learningRate = learningRate;
45+
this.maxEpochs = maxEpochs;
46+
}
47+
48+
/**
49+
* Fits the classifier using binary training labels.
50+
*
51+
* <p>Fitting resets the weights and bias to zero before training. The
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.
57+
*
58+
* @param features training feature vectors
59+
* @param labels corresponding binary labels, each either {@code 0} or
60+
* {@code 1}
61+
* @throws IllegalArgumentException if the training data is invalid
62+
* @throws ArithmeticException if training diverges and the learned
63+
* parameters stop being finite
64+
*/
65+
public void fit(double[][] features, int[] labels) {
66+
validateTrainingData(features, labels);
67+
68+
weights = new double[features[0].length];
69+
bias = 0.0;
70+
epochsRun = 0;
71+
converged = false;
72+
73+
for (int epoch = 0; epoch < maxEpochs; epoch++) {
74+
boolean updated = false;
75+
76+
for (int sampleIndex = 0; sampleIndex < features.length; sampleIndex++) {
77+
int prediction = rawPredict(features[sampleIndex]);
78+
int error = labels[sampleIndex] - prediction;
79+
80+
if (error != 0) {
81+
update(features[sampleIndex], error);
82+
updated = true;
83+
}
84+
}
85+
86+
epochsRun = epoch + 1;
87+
if (!updated) {
88+
converged = true;
89+
break;
90+
}
91+
}
92+
93+
ensureParametersAreFinite();
94+
}
95+
96+
/**
97+
* Predicts the binary label for one sample.
98+
*
99+
* @param sample feature vector to classify
100+
* @return {@code 0} or {@code 1}
101+
* @throws IllegalStateException if the classifier has not been fitted
102+
* @throws IllegalArgumentException if the sample is invalid
103+
*/
104+
public int predict(double[] sample) {
105+
ensureFitted();
106+
validateSample(sample);
107+
return rawPredict(sample);
108+
}
109+
110+
/**
111+
* Predicts binary labels for a batch of samples.
112+
*
113+
* @param samples feature vectors to classify
114+
* @return one prediction for each sample
115+
* @throws IllegalStateException if the classifier has not been fitted
116+
* @throws IllegalArgumentException if the batch or one of its samples is
117+
* invalid
118+
*/
119+
public int[] predict(double[][] samples) {
120+
ensureFitted();
121+
if (samples == null) {
122+
throw new IllegalArgumentException("samples cannot be null");
123+
}
124+
125+
int[] predictions = new int[samples.length];
126+
for (int sampleIndex = 0; sampleIndex < samples.length; sampleIndex++) {
127+
predictions[sampleIndex] = predict(samples[sampleIndex]);
128+
}
129+
return predictions;
130+
}
131+
132+
/**
133+
* Returns a defensive copy of the learned feature weights.
134+
*
135+
* @return learned weights in feature order
136+
* @throws IllegalStateException if the classifier has not been fitted
137+
*/
138+
public double[] getWeights() {
139+
ensureFitted();
140+
return weights.clone();
141+
}
142+
143+
/**
144+
* Returns the learned bias term.
145+
*
146+
* @return learned bias
147+
* @throws IllegalStateException if the classifier has not been fitted
148+
*/
149+
public double getBias() {
150+
ensureFitted();
151+
return bias;
152+
}
153+
154+
/**
155+
* Reports whether training completed with an update-free epoch.
156+
*
157+
* @return {@code true} if an epoch completed without an update
158+
* @throws IllegalStateException if the classifier has not been fitted
159+
*/
160+
public boolean hasConverged() {
161+
ensureFitted();
162+
return converged;
163+
}
164+
165+
/**
166+
* Returns the number of epochs performed by the last fit.
167+
*
168+
* @return number of completed epochs
169+
* @throws IllegalStateException if the classifier has not been fitted
170+
*/
171+
public int getEpochsRun() {
172+
ensureFitted();
173+
return epochsRun;
174+
}
175+
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+
184+
private void update(double[] sample, int error) {
185+
for (int featureIndex = 0; featureIndex < weights.length; featureIndex++) {
186+
weights[featureIndex] += learningRate * error * sample[featureIndex];
187+
}
188+
bias += learningRate * error;
189+
}
190+
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+
198+
private void ensureFitted() {
199+
if (weights == null) {
200+
throw new IllegalStateException("classifier has not been fitted");
201+
}
202+
}
203+
204+
private void validateTrainingData(double[][] features, int[] labels) {
205+
if (features == null || labels == null) {
206+
throw new IllegalArgumentException("features and labels cannot be null");
207+
}
208+
if (features.length == 0 || labels.length == 0) {
209+
throw new IllegalArgumentException("features and labels cannot be empty");
210+
}
211+
if (features.length != labels.length) {
212+
throw new IllegalArgumentException("features and labels must have the same length");
213+
}
214+
if (features[0] == null || features[0].length == 0) {
215+
throw new IllegalArgumentException("feature vectors cannot be null or empty");
216+
}
217+
218+
int featureCount = features[0].length;
219+
for (int sampleIndex = 0; sampleIndex < features.length; sampleIndex++) {
220+
double[] sample = features[sampleIndex];
221+
if (sample == null) {
222+
throw new IllegalArgumentException("feature vectors cannot be null or empty");
223+
}
224+
if (sample.length != featureCount) {
225+
throw new IllegalArgumentException("all feature vectors must have the same dimension");
226+
}
227+
validateFiniteValues(sample);
228+
if (labels[sampleIndex] != 0 && labels[sampleIndex] != 1) {
229+
throw new IllegalArgumentException("labels must be either 0 or 1");
230+
}
231+
}
232+
}
233+
234+
private void validateSample(double[] sample) {
235+
if (sample == null || sample.length != weights.length) {
236+
throw new IllegalArgumentException("sample must match the training feature dimension");
237+
}
238+
validateFiniteValues(sample);
239+
}
240+
241+
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) {
248+
for (double value : values) {
249+
if (!Double.isFinite(value)) {
250+
return false;
251+
}
252+
}
253+
return true;
254+
}
255+
}

0 commit comments

Comments
 (0)