//Jack Carrozzo (jack {@} crepinc.com)
//gcc -o neural-backprop_g neural-backprop_g.c -lm -Wall -O3

/* backpropogation algorithm:
1) starting at output layer, recursivly compute gradient omega for each node.
  on an output node:  omega = a'(inp)*error = inp*(1-inp)*(target-outp)
  on some other node: omega = a'(inp)*SUM(w1*omega1...wN*omegaN) = inp*(1-inp)*(w1*omega1...)
    weights and omega in the sumk is from the next layer down

2) change weights according to:
  Dw=mu*omega*inp
*/

#include <stdio.h>
#include <stdlib.h>
#include <math.h>
#include <unistd.h>

#define NUM_IN   2
#define NUM_HID1 8
#define NUM_HID2 8
#define NUM_OUT  1


float mu=0.2;                         //learning rate
float mc=0.0;                         //momentum constant

float in[NUM_IN];                     //the two inputs (output values of nodes)
float hid1[NUM_HID1];                 //out values of the hidden nodes
float hid1_o[NUM_HID1];
#ifdef NUM_HID2
float hid2[NUM_HID2];
float hid2_o[NUM_HID2];
#endif
float out[NUM_OUT];
float out_o[NUM_OUT];
float target[NUM_OUT];

#ifdef NUM_HID2
float wa_in_hid1[NUM_IN][NUM_HID1];   //weights
float wa_hid1_hid2[NUM_HID1][NUM_HID2];
float wa_hid2_out[NUM_HID2][NUM_OUT];
float din[NUM_IN];
float dhid1[NUM_HID1];               //changes
float dhid2[NUM_HID2];
float dout[NUM_OUT];
float dwa_in_hid1[NUM_IN][NUM_HID1];   //weight changes
float dwa_hid1_hid2[NUM_HID1][NUM_HID2];
float dwa_hid2_out[NUM_HID2][NUM_OUT];
float dowa_in_hid1[NUM_IN][NUM_HID1];   //previous weight changes (for momentum)
float dowa_hid1_hid2[NUM_HID1][NUM_HID2];
float dowa_hid2_out[NUM_HID2][NUM_OUT];
#else
float wa_in_hid1[NUM_IN][NUM_HID1];     //weights
float wa_hid1_out[NUM_HID1][NUM_OUT];
float din[NUM_IN];
float dhid1[NUM_HID1];                  //changes
float dout[NUM_OUT];
float dwa_in_hid1[NUM_IN][NUM_HID1];    //weight changes
float dwa_hid1_out[NUM_HID1][NUM_OUT];
float dowa_in_hid1[NUM_IN][NUM_HID1];   //privious weight changes
float dowa_hid1_out[NUM_HID1][NUM_OUT];
#endif

pid_t getpid(void);
float sigmoid(float);
float normrand(void);
float floatrand(void);
void randweights(void);
void teach(int);
void back_prop(void);
void run_net(void);

int main( int argc, char *argv[] ) {

  randweights();

  int i;
  for (i=0;i<1000;i++) {
    in[0]=floatrand(); 
    in[1]=floatrand();
    run_net();
    target[0]=in[0]*in[1]; // teach multiplication
    back_prop();
    printf("%2.3f , %2.3f --> %2.3f (real: %2.3f error: %2.3f)\n",in[0],in[1],out[0],target[0],target[0]-out[0]);
  }
 
  return 0;
}

//teach: run the net, and control inputs to net as well
//as the back propogation function.
//in: max iterations (int)
void teach ( int iters ) {
  //int i;
  //float q;

  /*for (i=0;i<iters;i++) {
    q=(((float)(rand())/((float)(RAND_MAX))));
    
    if (q<0.25) {   
      result=run_net(0,0);
      back_prop(0.0);
    } else if (q<0.5) {
      result=run_net(1,0);
      back_prop(1.0);
    } else if (q<0.75) {
      result=run_net(0,1);
      back_prop(1.0);
    } else {
      result=run_net(1,1);
      back_prop(0.0);
    }
  }*/
  
}

//back_prop: adjust the weights using the back propogation method
void back_prop( void ) {

  int i,j;
  float omega_out[NUM_OUT];
  #ifdef NUM_HID2
  float omega_hid2[NUM_HID2];
  #endif
  float omega_hid1[NUM_HID1];
  float tmp;

  #ifdef NUM_HID2
    //evaluate omega for all nodes besides input

    for (i=0;i<NUM_OUT;i++) { //output nodes
      omega_out[i]=out[i]*(1.0-out[i])*(target[i]-out_o[i]);
    }

    for (i=0;i<NUM_HID2;i++) { //second hidden layer
      tmp=0.0;
      for (j=0;j<NUM_OUT;j++) {
        tmp+=wa_hid2_out[i][j]*omega_out[j];
      }
      omega_hid2[i]=hid2[i]*(1.0-hid2[i])*tmp;
    }

    for (i=0;i<NUM_HID1;i++) { //first hidden layer
      tmp=0.0;
      for (j=0;j<NUM_HID2;j++) {
        tmp+=wa_hid1_hid2[i][j]*omega_hid2[j];
      }
      omega_hid1[i]=hid1[i]*(1.0-hid1[i])*tmp;
    }

    //now change the weights

    for (i=0;i<NUM_HID2;i++) {
      for (j=0;j<NUM_OUT;j++) {
        dwa_hid2_out[i][j]=(mu*omega_out[j]*out[j])+(dowa_hid2_out[i][j]*mc);
      }
    }

    for (i=0;i<NUM_HID1;i++) {
      for (j=0;j<NUM_HID2;j++) {
        dwa_hid1_hid2[i][j]=(mu*omega_hid2[j]*hid2[j])+(dowa_hid1_hid2[i][j]*mc);
      }
    }

    for (i=0;i<NUM_IN;i++) {
      for (j=0;j<NUM_HID1;j++) {
        dwa_in_hid1[i][j]=(mu*omega_hid1[j]*hid1[j])+(dowa_in_hid1[i][j]*mc);
      }
    }

    //change weights from deltas, and save to old weights

    for (i=0;i<NUM_HID2;i++) {
      for (j=0;j<NUM_OUT;j++) {
        wa_hid2_out[i][j]+=dwa_hid2_out[i][j];
        dowa_hid2_out[i][j]=dwa_hid2_out[i][j];
      }
    }

    for (i=0;i<NUM_HID1;i++) {
      for (j=0;j<NUM_HID2;j++) {
        wa_hid1_hid2[i][j]+=dwa_hid1_hid2[i][j];
        dowa_hid1_hid2[i][j]=dwa_hid1_hid2[i][j];
      }
    }

    for (i=0;i<NUM_IN;i++) {
      for (j=0;j<NUM_HID1;j++) {
        wa_in_hid1[i][j]+=dwa_in_hid1[i][j];
        dowa_in_hid1[i][j]=dwa_in_hid1[i][j];
      }
    }
  #else
    for (i=0;i<NUM_OUT;i++) { //output nodes
      omega_out[i]=out[i]*(1.0-out[i])*(target[i]-out_o[i]);
    }

    for (i=0;i<NUM_HID1;i++) { //hidden 
      tmp=0.0;
      for (j=0;j<NUM_OUT;j++) {
        tmp+=wa_hid1_out[i][j]*omega_out[j];
      }
      omega_hid1[i]=hid1[i]*(1.0-hid1[i])*tmp;
    }

    //now change the weights

    for (i=0;i<NUM_HID1;i++) {
      for (j=0;j<NUM_OUT;j++) {
        dwa_hid1_out[i][j]=(mu*omega_out[j]*out[j])+(dowa_hid1_out[i][j]*mc);
      }
    }

    for (i=0;i<NUM_IN;i++) {
      for (j=0;j<NUM_HID1;j++) {
        dwa_in_hid1[i][j]=(mu*omega_hid1[j]*hid1[j])+(dowa_in_hid1[i][j]*mc);
      }
    }
  
    //change weights from deltas, and save to old weights

    for (i=0;i<NUM_HID1;i++) {
      for (j=0;j<NUM_OUT;j++) {
        wa_hid1_out[i][j]+=dwa_hid1_out[i][j];
        dowa_hid1_out[i][j]=dwa_hid1_out[i][j];
      }
    }

    for (i=0;i<NUM_IN;i++) {
      for (j=0;j<NUM_HID1;j++) {
        wa_in_hid1[i][j]+=dwa_in_hid1[i][j];
        dowa_in_hid1[i][j]=dwa_in_hid1[i][j];
      }
    }
  #endif
}

//run_net: hopefully this is self explanatory...
void run_net ( void ) { 
  int i,j;
  float tmp;

  //the array in[] needs to have been set manually before calling this function

  for (i=0;i<NUM_HID1;i++) {
    tmp=0.0;
    for (j=0;j<NUM_IN;j++) {
      tmp+=in[j]*wa_in_hid1[j][i];
    }
    hid1[i]=tmp;
    hid1_o[i]=sigmoid(hid1[i]);
  }

  #ifdef NUM_HID2
    for (i=0;i<NUM_HID2;i++) {
      tmp=0.0;
      for (j=0;j<NUM_HID1;j++) {
        tmp+=hid1[j]*wa_hid1_hid2[j][i];
      }
      hid2[i]=tmp;
      hid2_o[i]=sigmoid(hid2[i]);
    }

    for (i=0;i<NUM_OUT;i++) {
      tmp=0.0;
      for (j=0;j<NUM_HID2;j++) {
        tmp+=hid2[j]*wa_hid2_out[j][i];
      } 
      out[i]=tmp;
      out_o[i]=sigmoid(out[i]);
    }
  #else
    for (i=0;i<NUM_OUT;i++) {
      tmp=0.0;
      for (j=0;j<NUM_HID1;j++) {
        tmp+=hid1[j]*wa_hid1_out[j][i];
      }
      out[i]=tmp;
      out_o[i]=sigmoid(out[i]);
    }
  #endif

  //the output has been stored in out_o[]
}

//randweights: set the weights for the first run
void randweights ( void ) {
  int i,j;

  srand(getpid());

  #ifdef NUM_HID2
   for (i=0;i<NUM_IN;i++) {
     for (j=0;j<NUM_HID1;j++) {
       wa_in_hid1[i][j]=normrand();
     }
   }

   for (i=0;i<NUM_HID1;i++) {
     for (j=0;j<NUM_HID2;j++) {
       wa_hid1_hid2[i][j]=normrand();
     }
   }

   for (i=0;i<NUM_HID2;i++) {
     for (j=0;j<NUM_OUT;j++) {
       wa_hid2_out[i][j]=normrand();
     }
   }
  #else
   for (i=0;i<NUM_IN;i++) {
     for (j=0;j<NUM_HID1;j++) {
       wa_in_hid1[i][j]=normrand();
     }
   }

   for (i=0;i<NUM_HID1;i++) {
     for (j=0;j<NUM_OUT;j++) {
       wa_hid1_out[i][j]=normrand();
     }
   }
  #endif
}

//sigmoid: calculate the sigmoid function value
//regulates between 0 and 1 on an S curve
//in: the value to be calculated
float sigmoid ( float x ) {
  return (1/(1+exp(-x)));
}

//normrand: return a random number on the 
//interval [-1,1]
float normrand( void ) {
  return ((2.0*(((float)(rand()))/((float)(RAND_MAX))))-1.0);
}

//floatrand: 0-->1
float floatrand( void ) {
  float r=normrand();

  if (r>0.0) {
    return r;
  } else {
    return -r;
  }
}
