// gcc my-chol.c -o chol.x -Wall -L/Users/pauldj/works/libs/OpenBLAS/ -lopenblas 
// gcc my-chol.c -o chol.x -Wall -L/opt/local/lib/ -lopenblas


#include <stdlib.h> 
#include <stdio.h>  
#include <time.h>     // for time()
#include <math.h>     // for sqrt()

#define min(a, b) (((a) < (b)) ? (a) : (b))

#define get_ticks(var) { \
   unsigned int __a, __d; \
   asm volatile("rdtsc" : "=a" (__a), "=d" (__d)); \
   var = ((unsigned long) __a) | (((unsigned long) __d) << 32); \
   } while(0)                                                                     
                                                                                 
int dtrsm_(char *, char *, char *, char *, int *, int *,
	   double *, double *, int *, double *, int *);

int dgemm_(char *, char *, int *, int *, int *, double *,
	   double *, int *, double *, int *, double *, double *, int *);

int dsyrk_(char *, char *, int *, int *, double *, double *, int *,
	   double *, double *, int *);

void chol1( double *A, int n, int ldA, int b );
void mat_print( double *A, int n, int m, int ldA, char *text);


int main(int argc, char *argv[] )
{
   int n,i,j,b;
   double *A; //, *L;  
   double tmp;
   unsigned long ticks, ticksB4, ticksAFT;                                       
  
   if(argc <= 2) { printf("Arguments needed!\n"); return(-1); }
   if(argc > 2) {
      n = atoi(argv[1]);
      b = atoi(argv[2]);
   }

   srand48( (unsigned)time((time_t *)NULL) ); 

   A = (double *) malloc( n*n * sizeof(double) );

   for( i=0; i<n; i++ ) // row idx
      for( j=0; j<=i; j++ ) //col idx
	 {
	    tmp = drand48(); 
	    A[i+j*n] = tmp;
	    A[j+i*n] = tmp;
	    if(i==j) A[i+i*n] += n;
	 }

   //   mat_print( A, n, n, n, "A" );

   get_ticks(ticksB4);
   chol1( A, n, n, b );
   get_ticks(ticksAFT);

   //   mat_print( A, n, n, n, "L" );

   ticks = ticksAFT - ticksB4;
   printf(" cycles= %.5ld\n", ticks );
   
   return(0);
}


void mat_print( double *A, int n, int m, int ldA, char *text)
{
   int i,j;
   FILE *output;
   output = stdout;

   fprintf(output, "%s = [ ...\n", text );
   for( i=0; i<n; i++ ) //row
      {
	 for( j=0; j<m; j++ )  //column
	    fprintf(output, "%.15g ", A[i+j*ldA] );
	 fprintf(output, "; ...\n" );
      }
   fprintf(output, "];\n" );

   return;
}


void chol1( double *A, int n, int ldA, int b )
{
   int idx=0, blk=0;
   double alpha=1, beta=1;

   if( n == 1 ){ A[0] = sqrt(A[0]); return; }

   while( idx < n )
      {
	 blk = min( b, n-idx );
	 
	 if( idx > 0 )
	    {
	       // TRSM
	       alpha=1;
	       dtrsm_( "R", "L", "T", "N", 
		       &blk, &idx, &alpha, &A[0], &ldA, &A[idx], &ldA);
	       //mat_print( &A[ idx ], blk, idx, ldA, "A10" );
	       
	       // SYRK, GEMM
	       alpha=-1;
	       //dgemm_( "N", "T", &blk, &blk, &idx, &alpha, &A[idx], &ldA,
	       //        &A[idx], &ldA, &beta, &A[ idx + idx * ldA ], &ldA );

	       dsyrk_( "L", "N", &blk, &idx, &alpha, &A[idx], &ldA,
	               &beta, &A[ idx + idx * ldA ], &ldA );

	       //mat_print( &A[ idx + idx * ldA ], blk, blk, ldA, "A11" );

	    };

	 chol1( &A[ idx + idx * ldA ], blk, ldA, 1 );

	 idx = idx + blk;
      }

   return;
}
