sandeshMC

BackPropagation.java

Apr 15th, 2016
50
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Java 4.33 KB | None | 0 0
  1. package backpropagation;
  2.  
  3. import java.util.Arrays;
  4.  
  5. class Sigmoid {
  6.  
  7.     public static double output(double x) {
  8.         return 1.0 / (1.0 + Math.exp(-x));
  9.     }
  10.  
  11.     public static double derivative(double x) {
  12.         return x * (1 - x);
  13.     }
  14. }
  15.  
  16. class Neuron {
  17.  
  18.     Neuron(int n) {
  19.  
  20.         inputs = new double[n];
  21.         weights = new double[n];
  22.     }
  23.  
  24.     Neuron() {
  25.  
  26.     }
  27.     double inputs[];
  28.     double weights[];
  29.     double net = 0;
  30.     double bias;
  31.     double out;
  32.     double error;
  33.  
  34.     void calculateNet() {
  35.  
  36.         for (int i = 0; i < inputs.length; i++) {
  37.             net += inputs[i] * weights[i];
  38.         }
  39.         net += bias;
  40.     }
  41.  
  42.     void calculateOut() {
  43.  
  44.         out = Sigmoid.output(net);
  45.     }
  46.  
  47.     void randomizeWeights() {
  48.  
  49.         for (int i = 0; i < weights.length; i++) {
  50.             weights[i] = Math.random();
  51.         }
  52.         bias = Math.random();
  53.     }
  54.    
  55.     public void adjustWeights() {
  56.         for (int i = 0; i < weights.length; i++) {
  57.             double ed = 0;
  58.             for (int j = 0; j < inputs.length; j++) {
  59.                 ed += error * inputs[i];
  60.             }
  61.             weights[i] += ed * 0.1;
  62.         }
  63.         // weights[0] += (error * inputs[0]);
  64.         // weights[1] += (error * inputs[1]);
  65.         bias += error;
  66.     }
  67.  
  68. }
  69.  
  70. public class BackPropagation {
  71.  
  72.     public static double learningRate = 0.01;
  73.    
  74.    
  75.    
  76.  
  77.     public static void main(String[] args) {
  78.         double inputs[][] = {{0, 0}, {0, 1}, {1, 0}, {1, 1}};
  79.         double[] results = {0, 1, 1, 0};
  80.         // TODO code application logic here
  81.         Neuron hiddenNeurons[] = new Neuron[2];
  82.         hiddenNeurons[0] = new Neuron(2);
  83.         hiddenNeurons[1] = new Neuron(2);
  84.         Neuron outputNeuron = new Neuron(2);
  85.         outputNeuron.randomizeWeights();
  86.         for (int i = 0; i < hiddenNeurons.length; i++) {
  87.             hiddenNeurons[i].randomizeWeights();
  88.             System.out.println("Weights for " + i + " is " + Arrays.toString(hiddenNeurons[i].weights));
  89.         }
  90.         int epoch = 0;
  91.         while (epoch++ != 1000) {
  92.           //  System.out.println(" Inside while ");
  93.             for (int i = 0; i < 4; i++) {
  94.                
  95.                // System.out.println(" Inside for ");
  96.                 for (int j = 0; j < hiddenNeurons.length; j++) {
  97.                    
  98.                     hiddenNeurons[j].inputs = inputs[i];
  99.                     hiddenNeurons[j].calculateNet();
  100.                     hiddenNeurons[j].calculateOut();
  101.                    // System.out.println(" Inside hidden forward ");
  102.                 }
  103.                 for (int j = 0; j < outputNeuron.inputs.length; j++) {
  104.                     outputNeuron.inputs[j] = hiddenNeurons[j].out;
  105.                    // System.out.println(" Inside output forward ");
  106.                    
  107.                 }
  108.                 outputNeuron.calculateNet();
  109.                 outputNeuron.calculateOut();
  110.                 System.out.println("Epoch No. " + epoch);
  111.                 System.out.println("Output for " + i + " is " + outputNeuron.out);
  112.                 //System.out.println("Output for " + i + " is " + Arrays.toString(outputNeuron.weights));
  113.  
  114.                 outputNeuron.error = (outputNeuron.out - results[i]) * Sigmoid.derivative(outputNeuron.out);
  115.               /*  for (int j = 0; j < outputNeuron.weights.length; j++) {
  116.                     outputNeuron.weights[j] -= outputNeuron.error * hiddenNeurons[j].out * learningRate;
  117.                 }
  118.                 outputNeuron.bias += outputNeuron.error;*/
  119.  
  120.                 for (int j = 0; j < hiddenNeurons.length; j++) {
  121.                     hiddenNeurons[j].error = outputNeuron.error * outputNeuron.weights[j] * Sigmoid.derivative(hiddenNeurons[j].out);
  122.                     /*for (int k = 0; k < hiddenNeurons[j].weights.length; k++) {
  123.                         hiddenNeurons[j].weights[k] -= hiddenNeurons[j].error * hiddenNeurons[j].inputs[k] *learningRate;
  124.                     }
  125.                     hiddenNeurons[j].bias += hiddenNeurons[j].error;
  126.                    // System.out.println("Hidden layer for " + i + " is " + Arrays.toString(hiddenNeurons[j].weights));
  127.                     System.out.println("Weights for " + j + " is " + Arrays.toString(hiddenNeurons[j].weights));*/
  128.                 }
  129.  
  130.             }
  131.  
  132.         }
  133.  
  134.     }
  135.  
  136. }
Add Comment
Please, Sign In to add comment