|
| 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