/**************************************************************************
* Rete neurale modello Error Back Propagation                             *
* Esempio: Lancio sasso                                                   *
* Creato by G. Ligorio                                                    *
**************************************************************************/

/* Inclusione file di supporto */
#include <exec/types.h>
#include <libraries/dosextens.h>
#include <stdio.h>
#include <math.h>

/* Definizione massime unita' x, h ed y e massimo numero esempi. */
#define MAXX    51
#define MAXH    51
#define MAXY    50
#define MAXE    200

/* Definizione macro della funzione di trasferimento. */
#define T(A)    (1.0/(1.0+exp(-A)))

/* Definizione variabili di utilizzo generale. */
int Nx,Nh,Ny,Ne;
double Epsilon,ErrRete;
int i,k,j,p;
double AltMax,GitMax,TemMax;

/* Definizione vettori X, H, Y, D e vettori dei pesi W1 e W2. */
double X[MAXX],H[MAXH],Y[MAXY],D[MAXY];
double W1[MAXX][MAXH];
double W2[MAXH][MAXY];

/* Definizione struttura contenente l'esempio e relativo vettore. */
struct Esempio
{
  double Vel,Ang,Git,Tem;
} E[MAXE];

ULONG seed=363747373;

/* Definizione prototipi di funzione. */
ULONG Random(ULONG);
void BackPropagate(void);
void Esegui(void);
void Traccia(double,double);
void GeneraEsempi(void);
void Apprendi(void);
void Modifica(void);
void Visualizza(void);
void DefinisciLancio(void);
void Salva(void);
void Carica(void);

/* Generatore di numeri pseudo-casuali. */
ULONG Random(max)
register ULONG max;
{
  register ULONG s = seed;
    
  seed = ((s << 16) + s + s + s) & 0x7FFFFFFF;
  return(seed % max);
}

/* Procedura di calcolo passo da sommare ai pesi mediante la propagazione dell'errore all'indietro. */
void BackPropagate()
{
  register double delta,delep;
  register int k,j,i;
  double ErrH[MAXH];
  
  for (k=1; k<=Nh+1; k++) ErrH[k]=0.0;
  ErrRete=0.0;
  for (j=1; j<=Ny; j++)
  {
    if (abs(D[j]-Y[j]) > ErrRete) ErrRete = abs(D[j]-Y[j]);
    delta = (D[j]-Y[j])*Y[j]*(1.0-Y[j]);
    delep = delta*Epsilon;
    for (k=1; k<=Nh+1; k++)
    {
      ErrH[k] += delta*W2[k][j];
      W2[k][j] += delep*H[k];
    }
  }
  for (k=1; k<=Nh; k++)
  {
    delta = ErrH[k]*H[k]*(1.0-H[k]);
    delep = delta*Epsilon;
    for (i=1; i<=Nx+1; i++) W1[i][k] += delep*X[i];
  }
  return;
}

/* Procedura che esegue la rete. */
void Esegui()
{
  register double A;
  register int k,j,i;
    
  for (k=1; k<=Nh; k++) H[k]=0.0;
  for (j=1; j<=Ny; j++) Y[j]=0.0;
  for (k=1; k<=Nh; k++)
  {
    A=0.0;
    for (i=1; i<=Nx+1; i++) A += (W1[i][k]*X[i]);
    H[k]=T(A);
  }
  for (j=1; j<=Ny; j++)
  {
    A=0.0;
    for (k=1; k<=Nh+1; k++) A += (W2[k][j]*H[k]);
    Y[j]=T(A);
  }
  return;
}

/* Procedura di stampa dell'esempio. */
void Traccia(Vel,Ang)
register double Vel,Ang;
{
  printf("Velocita' tiro - %ld\n",(LONG)Vel);
  printf("Angolazione tiro - %f\n",Ang);
  return;
}

/* Generazione degli esempi casuali. */
void GeneraEsempi()
{
  register int p;
  char op[10];

  printf("Generazione Esempi... ");
  for (p=1; p<=Ne; p++)
  {
    E[p].Vel=((double)Random(100)*(70.0/100.0))+30.0;
    E[p].Ang=(((double)Random(90)*(50.0/90.0))+30.0)*(3.1415926/180.0);
    E[p].Git=(E[p].Vel*E[p].Vel*sin(E[p].Ang))/9.8;
    E[p].Tem=(E[p].Vel*2.0*sin(E[p].Ang))/9.8;
    Traccia(E[p].Vel,E[p].Ang);
  }
  printf("\nPremere RETURN per continuare -> "); gets(op);
  return;
}

/* Procedura di apprendimento. */
void Apprendi()
{
  register double ErrEpoca,ErrAmm;
  register int Pausa,Epoca,p;
  char c[80];
  
  printf("Numero esempi (1 - %ld) : ",MAXE); gets(c); Ne = atoi(c);
  printf("Errore massimo ammesso (0.001 - 0.99) : "); gets(c); ErrAmm=atof(c);
  printf("Tasso di apprendimento epsilon (0.001 - 0.99) : "); gets(c); Epsilon=atof(c);
  printf("Pausa ogni quante epoche (0 - 10000) : "); gets(c); Pausa=atoi(c);
  GeneraEsempi();
  Epoca=0;
  do
  {
    ErrEpoca=0.0;
    for (p=1; p<=Ne; p++)
    {
      X[1]=E[p].Git/GitMax; X[2]=E[p].Tem/TemMax;
      D[1]=E[p].Vel/100.0; D[2]=E[p].Ang/(3.1415926/2.0);
      Esegui();
      BackPropagate();
      if (ErrRete > ErrEpoca) ErrEpoca=ErrRete;
    }
    Epoca++;
    printf("Epoca = %ld, Errore = %1.5f.\n",(LONG)Epoca,ErrEpoca);
    if (Epoca % Pausa == 0)
    {
      printf("Continuare (s/N) : "); gets(c);
      if (c[0] == 'n' || c[0] == 'N') ErrEpoca=-1;
    }
  } while (ErrEpoca>=ErrAmm);
  return;
}

/* Modifica della struttura neurale; utilizzare 2 unita' di input, 5 del livello intermedio,
   e 2 del livello d'uscita; non esiste un metodo preciso per calcolare il numero delle unita' del livello
   intermedio, nella maggiorparte dei casi basta uno piu' il doppio del numero maggiore di unita' tra i
   livelli di input o di output. */
   
void Modifica()
{
  register int i,k,j;
  char c[80];

  printf("Numero unita' in input (1 - %ld) : ",MAXX-1); gets(c); Nx = atoi(c);
  printf("Numero unita' nascoste (1 - %ld) : ",MAXH-1); gets(c); Nh = atoi(c);
  printf("Numero unita' in output (1 - %ld) : ",MAXY); gets(c); Ny = atoi(c);
  for (i=1; i<=Nx; i++) X[i]=0.0; X[Nx+1]=1.0;
  for (k=1; k<=Nx; k++) H[k]=0.0; H[Nh+1]=1.0;
  for (j=1; j<=Nx; j++) Y[j]=0.0;
  for (i=1; i<=Nx+1; i++)
    for (k=1; k<=Nh; k++)
      W1[i][k] = ((double)Random(100)-50.0)/100.0;
  for (k=1; k<=Nh+1; k++)
    for (j=1; j<=Ny; j++)
      W2[k][j] = ((double)Random(100)-50.0)/100.0;
  AltMax = 5000.0/9.8;
  GitMax = 10000.0/9.8;
  TemMax = 200.0/9.8;
  return;
}

/* Visualizza i valori delle unita' e dei pesi. */
void Visualizza()
{
  char c[80];
  
  printf("\nUnita' di input :\n");
  for (i=1; i<=Nx+1; i++) printf("%1.2f ",X[i]); printf("\n");
  printf("\nUnita' nascoste :\n");
  for (k=1; k<=Nh+1; k++) printf("%1.2f ",H[k]); printf("\n");
  printf("\nUnita' di output :\n");
  for (j=1; j<=Ny; j++) printf("%1.2f ",Y[j]); printf("\n");
  printf("\nVisualizzare matrice W1 (S/n) "); gets(c);
  if (c[0]=='s' || c[0]=='S')
  {
    printf("\n    ");
    for (k=1; k<=Nh; k++) printf("%02ldH   ",(LONG)k); printf("\n");
    for (i=1; i<=Nx+1; i++)
    {
      printf("%02ldX ",(LONG)i);
      for (k=1; k<=Nh; k++) printf("%2.2f ",W1[i][k]);
      printf("\n");
    }
  }
  printf("\nVisualizzare matrice W2 (S/n) "); gets(c);
  if (c[0]=='s' || c[0]=='S')
  {
    printf("\n    ");
    for (j=1; j<=Ny; j++) printf("%02ldY   ",(LONG)j); printf("\n");
    for (k=1; k<=Nh+1; k++)
    {
      printf("%02ldX ",(LONG)k);
      for (j=1; j<=Ny; j++) printf("%2.2f ",W2[k][j]);
      printf("\n");
    }
  }
  return;
}

/* Procedura di richiesta dei dati in ingresso. */
void DefinisciLancio()
{
  char c[80];
  
  printf("Posizione bersaglio (0.0 - 1.0) : "); gets(c); X[1]=atof(c);
  printf("Tempo di caduta (1 - %ld) : ",(LONG)TemMax); gets(c); X[2]=atof(c)/TemMax;
  return;
}

/* Procedura di salvataggio dei pesi. */
void Salva()
{
  register struct FileHandle *fHandle;
  register LONG l;
  char nome[40];
  
  printf("Nome file - "); gets(nome);
  if ((fHandle = (struct FileHandle *)Open(nome,MODE_NEWFILE)) != NULL)
  {
    l = Write(fHandle,&W1[0][0],MAXX*MAXH*sizeof(double));
    l = Write(fHandle,&W2[0][0],MAXH*MAXY*sizeof(double));
    Close(fHandle);
  }
  else
  {
    printf("Errore scrittura, file non creato.\n");
  }
  return;
}
  
/* Procedura di caricamento dei pesi precedentemente salvati. */
void Carica()
{
  register struct FileHandle *fHandle;
  register LONG l;
  char nome[40];
  
  printf("Nome file - "); gets(nome);
  if ((fHandle = (struct FileHandle *)Open(nome,MODE_OLDFILE)) != NULL)
  {
    l = Read(fHandle,&W1[0][0],MAXX*MAXH*sizeof(double));
    l = Read(fHandle,&W2[0][0],MAXH*MAXY*sizeof(double));
    Close(fHandle);
  }
  else
  {
    printf("Errore lettura, caricamento non effettuato.\n");
  }
  return;
}

/* Programma principale che gestisce il menu' e chiama le diverse procedure. */
void main()
{
  char c[80];
  
  Modifica();
  do
  {
    printf("\n  Menu' opzioni:\n");
    printf("  1. -> Apprendi.\n");
    printf("  2. -> Esegui.\n");
    printf("  3. -> Visualizza.\n");
    printf("  4. -> Modifica.\n");
    printf("  5. -> Salva.\n");
    printf("  6. -> Carica.\n");
    printf("  7. -> Fine.\n");
    printf("\n  Scelta : "); gets(c);
    if (c[0]=='1' || c[0]=='a' || c[0]=='A') Apprendi();
    else if (c[0]=='2' || c[0]=='e' || c[0]=='E')
    {
      DefinisciLancio(); Esegui(); Traccia(Y[1]*100.0,Y[2]*(3.1415926/2.0));
    }
    else if (c[0]=='3' || c[0]=='v' || c[0]=='V') Visualizza();
    else if (c[0]=='4' || c[0]=='m' || c[0]=='M') Modifica();
    else if (c[0]=='5' || c[0]=='s' || c[0]=='S') Salva();
    else if (c[0]=='6' || c[0]=='c' || c[0]=='C') Carica();
  } while(c[0]!='7' && c[0]!='f' && c[0]!='F');
}
