import java.applet.*;
import java.awt.*;
import java.util.*;

public class backprop extends Applet {
  Network net;
  Button addPats;
  Button train;
  Button stop;
  Button resetWeights;
  TextField alpha, mu;
  TextArea messages;

  Graph error;

  public void init() {
    int[] nodes = {2,1};
    String[] nodeTypes = {"HiddenNode", "OutputNode"}; 
    net = new Network(nodes, nodeTypes);     
    addPats      = new Button("Add Patterns");
    train        = new Button("Train");
    stop         = new Button("Stop");
    resetWeights = new Button("Reset Network");
    error = new Graph(100,100,Color.red);
    alpha = new TextField(".2", 5);
    mu    = new TextField(".2", 5);
    messages = new TextArea("Press train to train the network\nPress stop to stop the training\nPress Reset Network to start over\n\n",10,35);
    messages.setEditable(false);

    setLayout(new BorderLayout());
    Panel north = new Panel();
    Panel south = new Panel();
    Panel center = new Panel();
    Panel errorPanel = new Panel();
    errorPanel.setLayout(new BorderLayout());

    north.add(addPats);
    north.add(train);
    north.add(stop);
    north.add(resetWeights);
    errorPanel.add("North", new Label("Aggregate Error", Label.CENTER));
    errorPanel.add("Center", error);
    center.add(errorPanel);
    center.add(messages);
    south.add(new Label("Learning Constant:", Label.CENTER));
    south.add(alpha);
    south.add(new Label("Momentum Constant:", Label.CENTER));
    south.add(mu);
    add("North", north);
    add("Center", center);
    add("South", south);
    validate();
  }
  public void start() {
    validate();
  }
  public void stop() {
    net.stop();
  }
  
  public boolean action(Event event, Object arg) {
    if(event.target == addPats) {
      double[][] patsin= {{1,1}, {1,0}, {0, 1}, {0, 0}};
      double[][] patsout={{1}, {1}, {1}, {0}};
      int i;
      for(i=0; i<4; i++)
	net.addPattern(patsin[i], patsout[i]);
      return true;
    }
    if(event.target == train) {
      double a, m;
      int returnCode;
      try {
	a=Double.valueOf(alpha.getText()).doubleValue();
	m=Double.valueOf(mu.getText()).doubleValue();
	returnCode = net.train(0.1,a,m,error);
	if(returnCode == 1)
	  messages.appendText("Training started\n");
	else if(returnCode==-1)
	  messages.appendText("Training in progress\n");	  
	else if(returnCode==-2)
	  messages.appendText("No training patterns\n");
      } catch (NumberFormatException e) {
	messages.appendText("Alpha and Mu must be numbers\n");
      }
      return true;
    }
    if(event.target == stop) {
      net.stop();
      return true;
    }
    if(event.target == resetWeights) {
      net.resetWeights();
      return true;
    }
    return false;
  }
}

class Pattern {
  double[] input;
  double[] target;
  int numInputs, numOutputs;
 
  public Pattern(int in, int out) {
    numInputs = in;
    numOutputs = out;
    input = new double[numInputs];
    target = new double[numOutputs];
  }
  public double[] getInput() {
    return input;
  }
  public void setInput(double[] ins) {
    int i;
    for(i=0; i<numInputs; i++)
      input[i]=ins[i];
  }
  public double[] getTarget() {
    return target;
  }
  public void setTarget(double[] tars) {
    int i;
    for(i=0; i<numOutputs; i++)
      target[i]=tars[i];
  }
}

abstract class Node {
  public final static double RAND=0;

  double netInput;
  double activation;
  int incomingSignals;
  double[] weights;
  double[] deltaWeights;
  double[] prevDeltaWeights;
  double deltaError;
  double alpha;
  double mu;

  public void init(int is, double initial, double a, double m) {
    int i;

    alpha = a;
    mu=m;
    incomingSignals = is;
    weights = new double[incomingSignals];
    deltaWeights = new double[incomingSignals];
    prevDeltaWeights = new double[incomingSignals];
    setWeights(initial);
  }
  public void setWeights(double value) {
    int i;
    for(i=0; i<incomingSignals; i++) {
      deltaWeights[i]=0;
      if(value == RAND)
	weights[i]=(Math.random() * 2) - 1;
      else
	weights[i]=value;
    }  
  }
  public void calcNetInput(double[] inputs) {
    int i;
    netInput = 0;
    for(i=0; i<incomingSignals; i++) {
      netInput += weights[i] * inputs[i];
    }
  }
  public void calcActivation () {
    activation = transfer(netInput);
  }
  public void calcDeltaError (double inError) {
    deltaError=inError * dTransfer(netInput);
  }
  public void calcDeltaWeights(double[] activations) { 
    int i;
    for(i = 0; i<incomingSignals; i++) {
      prevDeltaWeights[i] = deltaWeights[i];
      deltaWeights[i] = alpha * deltaError * activations[i];
    }
  }
  public void updateWeights () {
    int i;
    for(i = 0; i<incomingSignals; i++) {
      weights[i] = weights[i] + deltaWeights[i] + mu*prevDeltaWeights[i];
    }
  }
  public double getActivation () {
    return activation;
  }
  public double getNetInput() {
    return netInput;
  }
  public double[] getWeights() {
    return weights;
  }
  public double getDeltaErrorForNode(int node) {
    return deltaError * weights[node];
  }
  abstract double transfer(double x);
  abstract double dTransfer(double x);
}

class HiddenNode extends Node {
  public HiddenNode() {
    incomingSignals = 0;
  }
  public HiddenNode(int is, double initial, double a, double m) {
    init(is, initial, a,m);
  }
  double transfer(double x) {
    return (1/(1+Math.exp(-x)));
  }
  double dTransfer(double x) {
    return (transfer(x) * (1-transfer(x)));
  }
}

class OutputNode extends Node {
  public OutputNode() {
    incomingSignals = 0;
  }
  public OutputNode(int is, double initial, double a, double m) {
    init(is, initial, a, m);
  }
  double transfer(double x) {
    return (1/(1+Math.exp(-x)));
  }
  double dTransfer(double x) {
    return (transfer(x) * (1-transfer(x)));
  }
}

class Layer {
  int numberOfNodes;
  Node[] nodes;
  double[] inputs;
  double[] outputs;
  
  public Layer() {
    numberOfNodes = 0;
  }
  public Layer(int nn, String nodeType, int is, double alpha, double mu) {
    init(nn, nodeType, is, alpha, mu);
  }
  public void init(int nn, String nodeType, int is, double alpha, double mu) {
	int i;
	
	numberOfNodes = nn;

	inputs = new double[is];
	outputs = new double[numberOfNodes];
	
	nodes = new Node[numberOfNodes];
	try {
	  for(i=0; i<nodes.length; i++) {
	    nodes[i] = (Node)Class.forName(nodeType).newInstance();
	    nodes[i].init(is, Node.RAND, alpha, mu);
	  }
	} catch (InstantiationException e) {
	  System.err.println("Can not create instance: "+e.getMessage());
	} catch (IllegalAccessException e) {
	  System.err.println("Illegal Access: "+e.getMessage());
	} catch (ClassNotFoundException e) {
	  System.err.println("Class not Found: "+e.getMessage());
	}
  }
  public void evaluate() {
    int i;
    for(i=0; i<numberOfNodes; i++) {
      nodes[i].calcNetInput(inputs);
      nodes[i].calcActivation();
      outputs[i]=nodes[i].getActivation();
    }
  }
  public void calcDeltaError(Layer above) {
    int i;
    for(i=0; i<numberOfNodes; i++) {
      nodes[i].calcDeltaError(above.getDeltaErrorForNode(i));
    }
  }
  public void calcDeltaError(double[] target) {
    int i;
    for(i=0; i<numberOfNodes; i++) {
      nodes[i].calcDeltaError(target[i]-outputs[i]);
    }
  }
  public void adjustWeights(double[] activations) {
    int i;
    for(i=0; i<numberOfNodes; i++) {
      nodes[i].calcDeltaWeights(activations);
      nodes[i].updateWeights();
    }
  }
  public double getDeltaErrorForNode(int n) {
    int i;
    double sum=0;

    for(i=0; i<numberOfNodes; i++) {
      sum+=nodes[i].getDeltaErrorForNode(n);
    }
    return sum;
  }
  public double[] getOutput() {
    return outputs;
  }

  public double[] getWeights(int n) {
    return nodes[n].getWeights();
  }
  public int getNumberOfNodes() {
    return numberOfNodes;
  }
  public void setInput(double[] ins) {
    int i;
    for(i=0; i<inputs.length; i++) {
      inputs[i]=ins[i];
    }
  }
  public void resetWeights() {
    int i;
    for(i=0; i<numberOfNodes; i++)
      nodes[i].setWeights(Node.RAND);
  }
}

class Network implements Runnable {
  Thread trainingThread = null;
  int numLayers, numPatterns;
  Layer[] layers;
  Vector trainingPatterns;
  double[] input, output;
  double alpha, mu, maxError;
  double aggregateError;
  int valid;
  Graph errorChart;

  public Network() {
    numLayers = 0;
    numPatterns = 0;
    alpha = .2;
    mu=.2;
    trainingPatterns = new Vector();
  }
  public Network(int[] nodes, String[] nodeType) {
    init(nodes, nodeType);
  }
  public void init(int[] nodes, String[] nodeType) {
    int i;
    numLayers = nodes.length;
    layers = new Layer[numLayers];
    numPatterns = 0;
    alpha = .2;
    mu = .2;
    for(i=0; i<numLayers; i++) {
      int is;
      if(i == 0)
	is = nodes[i];
      else
	is = layers[i-1].getNumberOfNodes();
      layers[i] = new Layer(nodes[i], nodeType[i], is ,alpha,mu);     
    }
    trainingPatterns = new Vector();
  }
  public void addPattern(double[] ins, double[] targets) {
    Pattern p = new Pattern(layers[0].getNumberOfNodes(), layers[numLayers-1].getNumberOfNodes());
    p.setInput(ins);
    p.setTarget(targets);
    trainingPatterns.addElement(p);
    numPatterns++;
  }
  public double[] evaluate(double[] in) {
    int i;
    if(numLayers > 0) {
      layers[0].setInput(in);
      layers[0].evaluate();
      for(i=1; i<numLayers; i++) {
	layers[i].setInput(layers[i-1].getOutput());
	layers[i].evaluate();
      }
      return layers[numLayers-1].getOutput();
    }
    return null;
  }
  public int train(double me, double a, double m, Graph error) {
    if(numPatterns == 0 || numLayers == 0)
      return -2;
    if(trainingThread == null || !trainingThread.isAlive()) {
      alpha = a;
      mu = m;
      maxError = me;
      errorChart = error;
      trainingThread = new Thread(this, "trainingThread");
      trainingThread.start();
      return 1;
    }
    else 
      return -1;
  }
  public void run() {
    Pattern p;
    int i,j,k=0;

    do {
      aggregateError = 0;
      for(j=0; j<numPatterns; j++) {
	p=(Pattern)trainingPatterns.elementAt(j);
	evaluate(p.getInput());
	layers[numLayers-1].calcDeltaError(p.getTarget());
	aggregateError += calcPatternError(p.getTarget());
	for(i=numLayers-2; i>=1; i--) {
	  layers[i].calcDeltaError(layers[i+1]);
	}
	layers[0].adjustWeights(p.getInput());
	for(i=1; i<numLayers; i++) {
	  layers[i].adjustWeights(layers[i-1].getOutput());
	}
      }
      k++;
      if(k>=100) {
	errorChart.addData(aggregateError);
	System.out.println(aggregateError);
	k=0;
      }
      try {
	trainingThread.sleep(6);
      } catch (InterruptedException e) {}
    }while(aggregateError > maxError);
    stop();
  }
  public void stop() {
    System.out.println("Training stopped");
    if(trainingThread != null)
      trainingThread.stop();
  }

  double calcPatternError(double[] target) {
    int i;
    double error;
    double[] output = layers[numLayers-1].getOutput();

    error = 0;
    for(i=0; i<target.length; i++)    
      error += Math.pow(target[i]-output[i],2);
    return error;
  }
  double getAggregateError() {
    return aggregateError;
  }
  public void resetWeights() {
    int i;
    for(i=0; i<numLayers; i++) {
      layers[i].resetWeights();
    }
  }
}

class Graph extends Canvas {
  Dimension size;
  Color graphColor;
  Vector data;

  public Graph(int width, int height, Color c) {
    size = new Dimension(width+2,height+2);
    graphColor = new Color(c.getRed(), c.getGreen(), c.getBlue());
    data = new Vector();
    int i;
    for(i=0; i<width; i++) {
      data.addElement(new Double(0));
    }
  } 

  public void addData(double d) {
    data.addElement(new Double(d*100));
    data.removeElementAt(0);
    repaint();
  }
  public void paint(Graphics g) {
    int i;
    Double d;
    g.setColor(Color.black);
    g.drawRect(0,0,size.width-1, size.height-1);
    g.setColor(graphColor);
    for(i=0; i<data.size(); i++) {
      d=(Double)data.elementAt(i);
      g.drawLine(i+1,size.height,i+1,size.height-d.intValue());
    }
  }
  public Dimension preferredSize() {
    return size;
  }
  public Dimension minimumSize() {
    return size;
  }
}







