#include <stdio.h>

float p[3][8][8], q[2][8][8], ca[3][8], cu[8], cx[2][7];
int opt[8][11];
FILE *fout;
main() {
  int i,j,n, pchoice, qchoice;
  float jtrue[11][8],fun(),y;
  void set_p_mtrx(),set_q_mtrx(), set_ca_mtrx(), set_cu_mtrx(), set_cx_mtrx(),
       init_to_zero(), pick_mtrx(), set_opt_mtrx(), fprint_mtrx(), iprint_mtrx();

  init_to_zero(p, 3);
  init_to_zero(q, 2);
  
  fout=fopen("/mit/smmadana/mdp.dat","w");
  pick_mtrx(&pchoice, &qchoice);

  set_p_mtrx(p,pchoice);
  set_q_mtrx(q,qchoice);
  set_ca_mtrx(ca);
  set_cu_mtrx(cu);
  set_cx_mtrx(cx);
  set_opt_mtrx(opt);

  for(j=0;j<8;j++) 
    jtrue[10][j]=0.0;

  for(n=9;n>=0;n--) 
    for(i=0;i<8;i++)
      jtrue[n][i]=fun(qchoice,n,i,jtrue);

  fprintf(fout,"pchoice=%d\n",pchoice);
  fprintf(fout,"qchoice=%d\n\n",qchoice);

  fprintf(fout,"                                Jtrue Array\n");
  fprint_mtrx(jtrue);
  fprintf(fout,"\n\n\n\n\n\n");
  fprintf(fout,"                                 Opt Array\n");
  iprint_mtrx(opt);
  fclose(fout);
}
/*****************************************************************************/
float fun (qchoice,n,i,jtrue)
int qchoice,n,i;
float jtrue[11][8]; {
  float cost=0,futcost,alpha=0.9524;
  int j,l;

  for(j=0;j<8;j++) {
    futcost=0.0;
    for(l=0;l<8;l++) 
      futcost+=p[opt[j][n]][i][l]*(cu[l]+jtrue[n+1][l]);
    cost+=q[1][i][j]*(futcost);
  }
  return(cost);
}
/*****************************************************************************/
void init_to_zero(a,i)
short i;
float a[][8][8];{
  short loop1, loop2, loop3;

  for(loop1=0;loop1<i;loop1++)  
    for(loop2=0;loop2<8;loop2++)
	for(loop3=0;loop3<8;loop3++)
	  a[loop1][loop2][loop3]=0.0;
}
/*****************************************************************************/
void pick_mtrx(pchoice, qchoice)
int *pchoice,*qchoice; {
  printf("\n\nEnter choices of matrixes...\n");
  printf("Choice of p matrix:");
  scanf("%d", pchoice);
  printf("Choice of q matrix:");
  scanf("%d", qchoice);
  printf("\n");
}
/*****************************************************************************/
void set_opt_mtrx(opt)
int opt[8][11]; {
  int n,j;

  for(j=0;j<8;j++)
    for(n=0;n<10;n++)
       if(j<4)
          opt[j][n]=2;
       else
        if(j<6)
          opt[j][n]=1;
       else
          opt[j][n]=0;
      
} 
/*****************************************************************************/
void fprint_mtrx(mtrx)
float mtrx[11][8]; {
  int i,j;

  fprintf(fout,"\n");
  fprintf(fout,"        ");
  for(j=0;j<8;j++)
    fprintf(fout," state %d ",j);
  fprintf(fout,"\n\n");
  for(i=0;i<10;i++) {
    fprintf(fout,"level %d",i);
    for(j=0;j<8;j++) {
      fprintf(fout," %8.4f",mtrx[i][j]);
    }
    fprintf(fout,"\n\n");
  }
  fprintf(fout,"\n\n");
} 
/*****************************************************************************/
void iprint_mtrx(mtrx)
int mtrx[8][11]; {
  int i,j;

  fprintf(fout,"\n");
  fprintf(fout,"        ");
  for(j=0;j<8;j++)
    fprintf(fout," state %d ",j);
  fprintf(fout,"\n\n");
  for(i=0;i<10;i++) {
    fprintf(fout,"level %d",i);
    for(j=0;j<8;j++) {
      fprintf(fout,"%6d   ",mtrx[j][i]);
    }
    fprintf(fout,"\n\n");
  }
}  
/*****************************************************************************/
#include "proc/set_mtrx.proc~~"
