/*     Copyright 2007 Francesco 'SkZ' Mauro */
/*     Any use and redistribution of this software must be authorized by the author*/


/*     This program is distributed in the hope that it will be useful, */
/*     but WITHOUT ANY WARRANTY; without even the implied warranty of */
/*     MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  */

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

typedef struct{
  long nr;
  long nc;
  long double **data;
  double det;
  int fl_det;

/*  long *row=NULL;
  long *col=NULL;*/

} SkZmatrix;
/***********************Mode*********************************/

SkZmatrix* SkZmatrix_calloc(long nr, long nc);
SkZmatrix* SkZmatrix_createfromarray(double **array, long nr, long nc);
SkZmatrix* SkZmatrix_clone(const SkZmatrix *matrix);
SkZmatrix* SkZmatrix_makeident(long nr);
void       SkZmatrix_free(SkZmatrix *matrix);

long SkZmatrix_print(const SkZmatrix *matrix, FILE *fout);

void    SkZmatrix_rowexchange(SkZmatrix *matrix, long this, long with_this);
void    SkZmatrix_colexchange(SkZmatrix *matrix, long this, long with_this);
long double* SkZmatrix_rowget(const SkZmatrix *matrix, long this);
long double* SkZmatrix_colget(const SkZmatrix *matrix, long this);

void SkZmatrix_rowaddrowmulby(SkZmatrix *matrix, long this, long add_this, double mul_by);
void SkZmatrix_rowaddrowdivby(SkZmatrix *matrix, long this, long add_this, double div_by);
void SkZmatrix_coladdcolmulby(SkZmatrix *matrix, long this, long add_this, double mul_by);
void SkZmatrix_coladdcoldivby(SkZmatrix *matrix, long this, long add_this, double div_by);

void SkZmatrix_mulby(SkZmatrix *matrix, double k);
void SkZmatrix_divby(SkZmatrix *matrix, double k);
void SkZmatrix_rowmul(SkZmatrix *matrix, long this, double k);
void SkZmatrix_colmul(SkZmatrix *matrix, long this, double k);
void SkZmatrix_rowdiv(SkZmatrix *matrix, long this, double k);
void SkZmatrix_coldiv(SkZmatrix *matrix, long this, double k);

SkZmatrix* SkZmatrix_detandinv(SkZmatrix *matrix);
long double*    SkZmatrix_mulbyvec(const SkZmatrix *matrix, const double *vector);
SkZmatrix* SkZmatrix_mul(const SkZmatrix *matrix1, const SkZmatrix *matrix2);
SkZmatrix* SkZmatrix_plus(const SkZmatrix *matrix1, const SkZmatrix *matrix2);
SkZmatrix* SkZmatrix_minus(const SkZmatrix *matrix1, const SkZmatrix *matrix2);


/********************************CODE*********************************/

SkZmatrix *SkZmatrix_calloc(long nr, long nc){
  SkZmatrix *matrix;
  long i;

  if(!(nr>0 && nc>0)) return NULL;
  
  matrix=(SkZmatrix *) malloc(sizeof(SkZmatrix));
  matrix->nr=nr, matrix->nc=nc;

  if(!(matrix->data=(long double **) calloc(matrix->nr, sizeof(long double *)))) return NULL; 
  for(i=0; i<matrix->nr; i++) 
    if(!(matrix->data[i]=(long double *) calloc(matrix->nc, sizeof(long double)))) return NULL;
  matrix->det=1;
  matrix->fl_det=0;
  
  return matrix;
}


SkZmatrix * SkZmatrix_createfromarray(double **array, long nr, long nc){
  long i,j;
  SkZmatrix *matrix;
  
  if((matrix=SkZmatrix_calloc(nr,nc))) return NULL;
  for(i=0; i<matrix->nr; i++) for(j=0; j<matrix->nc; j++) matrix->data[i][j]=array[i][j];

  return matrix; 
}

SkZmatrix* SkZmatrix_clone(const SkZmatrix *matrix){
  long i,j;
  SkZmatrix *clone;
  
  if(!(clone=SkZmatrix_calloc(matrix->nr, matrix->nc))) return NULL;
  for(i=0; i<clone->nr; i++)for(j=0; j<clone->nc; j++) clone->data[i][j]=matrix->data[i][j];
  clone->det=matrix->det, clone->fl_det=matrix->fl_det;

  return clone; 
}

SkZmatrix *SkZmatrix_makeident(long nr){
  SkZmatrix *matrix;
  long i;

  if(!(matrix=SkZmatrix_calloc(nr, nr)))return NULL;
  for(i=0; i<matrix->nr; i++) matrix->data[i][i]=1;
  matrix->det=1, matrix->fl_det=1;

  return matrix; 
}


void SkZmatrix_free(SkZmatrix *matrix){
  long i;
  
  if(matrix){
    for(i=0; i<matrix->nr; i++) 
      free(matrix->data[i]);
    free(matrix->data);
    free(matrix);
  }
}

/*-----------------------------------*/


long SkZmatrix_print(const SkZmatrix *matrix, FILE *fout){
  long i,j,tot,fl;
  
  if(matrix && fout)
    for(tot=fl=i=0; i<matrix->nr && fl>=0; i++, tot+=(fl=fprintf(fout,"\n")))
      for(j=0; j<matrix->nc && fl>=0; j++) 
	tot+=(fl=fprintf(fout,"\t%12Le", matrix->data[i][j]));  
  else  
    return -1;
  return   (fl>=0)? tot: -1;
}

/*-----------------------------------*/

void SkZmatrix_rowexchange(SkZmatrix *matrix, long this, long with_this){
  long double *temp;
  
  temp=matrix->data[with_this];
  matrix->data[with_this]=matrix->data[this];
  matrix->data[this]=temp;  
  matrix->det=-matrix->det;
}


void SkZmatrix_colexchange(SkZmatrix *matrix, long this, long with_this){
  long double *col;
  long i;
  
  col=(long double *) calloc(matrix->nc, sizeof(long double));
  for(i=0; i<matrix->nr; i++) col[i]=matrix->data[i][with_this];

  for(i=0; i<matrix->nr; i++) matrix->data[i][with_this]=matrix->data[i][this];
  for(i=0; i<matrix->nr; i++) matrix->data[i][this]=col[i];  
  matrix->det=-matrix->det;

  free(col);
}



long double* SkZmatrix_rowget(const SkZmatrix *matrix, long this){
  long double *row;
  long i;
  
  row=(long double *) calloc(matrix->nc, sizeof(long double));
  for(i=0; i<matrix->nc; i++) row[i]=matrix->data[this][i];

  return row;
}

long double* SkZmatrix_colget(const SkZmatrix *matrix, long this){
  long double *col;
  long i;
  
  col=(long double *) calloc(matrix->nc, sizeof(long double));
  for(i=0; i<matrix->nr; i++) col[i]=matrix->data[i][this];

  return col;
}


/*++++++++++++++++++++++++++++++++++++*/

void SkZmatrix_rowaddrowmulby(SkZmatrix *matrix, long this, long add_this, double mul_by){
  long i;

  
  for(i=0; i<matrix->nc; i++) matrix->data[this][i]+=matrix->data[add_this][i]*mul_by;
    
}


void SkZmatrix_rowaddrowdivby(SkZmatrix *matrix, long this, long add_this, double div_by){
  long i;

  for(i=0; i<matrix->nc; i++) matrix->data[this][i]+=matrix->data[add_this][i]/div_by;
    
}

void SkZmatrix_coladdcolmulby(SkZmatrix *matrix, long this, long add_this, double mul_by){
  long i;

  for(i=0; i<matrix->nc; i++) matrix->data[i][this]+=matrix->data[i][add_this]*mul_by;
    
}


void SkZmatrix_coladdcoldivby(SkZmatrix *matrix, long this, long add_this, double div_by){
  long i;

  for(i=0; i<matrix->nc; i++) matrix->data[i][this]+=matrix->data[i][add_this]/div_by;
    
}
/*%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%%*/

void SkZmatrix_mulby(SkZmatrix *matrix, double k){
  long i,j;
  
  for(i=0; i<matrix->nr; matrix->det*=k, i++) for(j=0; j<matrix->nc; matrix->data[i][j++]*=k);

}

void SkZmatrix_divby(SkZmatrix *matrix, double k){
  long i,j;
  double invk=1/k;
  
  for(i=0; i<matrix->nr; matrix->det*=invk, i++) for(j=0; j<matrix->nc; matrix->data[i][j++]*=invk);
}


void SkZmatrix_rowmul(SkZmatrix *matrix, long this, double k){
  long i;
  
  for(matrix->det*=k, i=0; i<matrix->nc; matrix->data[this][i++]*=k);  
}

void SkZmatrix_colmul(SkZmatrix *matrix, long this, double k){
  long i;
  
  for(matrix->det*=k, i=0; i<matrix->nr; matrix->data[i++][this]*=k);  
}


void SkZmatrix_rowdiv(SkZmatrix *matrix, long this, double k){
  long i;
  double invk=1/k;
  
  for(matrix->det*=invk, i=0; i<matrix->nc; matrix->data[this][i++]*=invk);  
}

void SkZmatrix_coldiv(SkZmatrix *matrix, long this, double k){
  long i;
  double invk=1/k;
  
  for(matrix->det*=invk, i=0; i<matrix->nr; matrix->data[i++][this]*=invk);  
}

/*&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&&*/

SkZmatrix * SkZmatrix_detandinv(SkZmatrix *matrix){
  SkZmatrix *inv, *clone;
  long i, j;
  double oldet, a;

  if(matrix->fl_det==1) if(!(oldet=matrix->det)) return NULL;;
  if(matrix->nr!=matrix->nc) return NULL;
  

  
  
  if(!( inv=SkZmatrix_makeident(matrix->nr))) return NULL;
  
  if(matrix->nr==1){
    matrix->det=matrix->data[0][0];
    matrix->fl_det=1;
    if(matrix->det){
      inv->det=inv->data[0][0]=1/matrix->data[0][0];
      inv->fl_det=1;
      return inv;
    }
    SkZmatrix_free(inv);
    return NULL;
  }

 
  
  if(matrix->nr==2){
    matrix->det=matrix->data[0][0]*matrix->data[1][1]-matrix->data[1][0]*matrix->data[0][1];
    matrix->fl_det=1;
    if(matrix->det){
      inv->data[0][0]=matrix->data[1][1]/matrix->det;
      inv->data[0][1]=-matrix->data[0][1]/matrix->det;
      inv->data[1][0]=-matrix->data[1][0]/matrix->det;
      inv->data[1][1]=matrix->data[0][0]/matrix->det;
      inv->det=1/matrix->det;
      inv->fl_det=1;
   
      return inv;
    }
    SkZmatrix_free(inv);
    return NULL;
  }


if(matrix->nr==3){
  matrix->det=matrix->data[0][0]*matrix->data[1][1]*matrix->data[2][2]+matrix->data[0][1]*matrix->data[1][2]*matrix->data[2][0]+matrix->data[0][2]*matrix->data[1][0]*matrix->data[2][1]-matrix->data[0][0]*matrix->data[1][2]*matrix->data[2][1]-matrix->data[0][1]*matrix->data[1][0]*matrix->data[2][2]-matrix->data[0][2]*matrix->data[1][1]*matrix->data[2][0];
    matrix->fl_det=1;
    if(matrix->det){
      inv->data[0][0]=(matrix->data[1][1]*matrix->data[2][2]-matrix->data[1][2]*matrix->data[2][1])/matrix->det;
      inv->data[0][1]=(matrix->data[0][2]*matrix->data[2][1]-matrix->data[0][1]*matrix->data[2][2])/matrix->det;
      inv->data[0][2]=(matrix->data[0][1]*matrix->data[1][2]-matrix->data[0][2]*matrix->data[1][1])/matrix->det;

      inv->data[1][0]=(matrix->data[1][2]*matrix->data[2][0]-matrix->data[1][0]*matrix->data[2][2])/matrix->det;
      inv->data[1][1]=(matrix->data[0][0]*matrix->data[2][2]-matrix->data[0][2]*matrix->data[2][0])/matrix->det;
      inv->data[1][2]=(matrix->data[0][2]*matrix->data[1][0]-matrix->data[0][0]*matrix->data[1][2])/matrix->det;

      inv->data[2][0]=(matrix->data[1][0]*matrix->data[2][1]-matrix->data[1][1]*matrix->data[2][0])/matrix->det;
      inv->data[2][1]=(matrix->data[0][1]*matrix->data[2][0]-matrix->data[0][0]*matrix->data[2][1])/matrix->det;
      inv->data[2][2]=(matrix->data[0][0]*matrix->data[1][1]-matrix->data[0][1]*matrix->data[1][0])/matrix->det;
      inv->det=1/matrix->det;
      inv->fl_det=1;
   
      return inv;
    }
    SkZmatrix_free(inv);
    return NULL;
  }
   

 
  if(!( clone=SkZmatrix_clone(matrix))) return NULL;
  
  /* Start */
  for(clone->det=clone->fl_det=1, i=0; i<clone->nr; i++){
    for(j=i; j<clone->nr && !clone->data[j][i]; j++);
    if (j==clone->nr) {
      matrix->det=0, matrix->fl_det=1; 
      SkZmatrix_free(inv);  SkZmatrix_free(clone); 
      return NULL;}
    if (j!=i) {
      SkZmatrix_rowexchange(clone, i, j); 
      SkZmatrix_rowexchange(inv, i, j);
    }

   

    if(clone->data[i][i]!=1){
      SkZmatrix_rowdiv(clone, i, a=clone->data[i][i]);
      SkZmatrix_rowdiv(inv, i, a);
    }
    
    for(j=i+1; j<clone->nr; j++)
      if(clone->data[j][i]!=0){
	SkZmatrix_rowaddrowmulby(clone, j, i, a=-(clone->data[j][i]));
	SkZmatrix_rowaddrowmulby(inv, j, i, a);
      }
  }
  

  for(i=clone->nr-2; i>=0; i--){
    for(j=clone->nr-1; j>i; j--)
      if(clone->data[i][j]!=0){
	SkZmatrix_rowaddrowmulby(clone, i, j, a=-(clone->data[i][j]));
	SkZmatrix_rowaddrowmulby(inv, i, j, a);
      }        
  }

  
  matrix->det=1/inv->det, matrix->fl_det=1;  
  
  return inv;
}


long double * SkZmatrix_mulbyvec(const SkZmatrix *matrix, const double *vector){
  long i, j;
  long double *out;
  
  out=(long double *) calloc(matrix->nr, sizeof(long double));
  
  for(i=0; i<matrix->nr; i++) 
    for(j=0; j<matrix->nc; j++) 
      out[i]+=matrix->data[i][j]*vector[j];
  
  
  return out;
}


SkZmatrix* SkZmatrix_mul(const SkZmatrix *matrix1, const SkZmatrix *matrix2){
  SkZmatrix *prod;
  long i, j, k;
  
  if(matrix1->nc!=matrix2->nr) return NULL;
  
  prod=SkZmatrix_calloc(matrix1->nr, matrix2->nc);
  
  for(i=0; i<prod->nr; i++) 
    for(j=0; j<prod->nc; j++) 
      for(k=0; k<matrix1->nc; k++)        
	  prod->data[i][j]+=matrix1->data[i][k]*matrix2->data[k][j];
  
  return prod;
}

SkZmatrix* SkZmatrix_plus(const SkZmatrix *matrix1, const SkZmatrix *matrix2){
  SkZmatrix *prod;
  long i, j;
  
  if(matrix1->nc!=matrix2->nc && matrix1->nr!=matrix2->nr) return NULL;
  
  prod=SkZmatrix_calloc(matrix1->nr, matrix1->nc);
  
  for(i=0; i<prod->nr; i++) 
    for(j=0; j<prod->nc; j++) 
      prod->data[i][j]=matrix1->data[i][j]+matrix2->data[i][j];
  
  return prod;
}


SkZmatrix* SkZmatrix_minus(const SkZmatrix *matrix1, const SkZmatrix *matrix2){
SkZmatrix *prod;
  long i, j;
  
  if(matrix1->nc!=matrix2->nc && matrix1->nr!=matrix2->nr) return NULL;
  
  prod=SkZmatrix_calloc(matrix1->nr, matrix1->nc);
  
  for(i=0; i<prod->nr; i++) 
    for(j=0; j<prod->nc; j++) 
      prod->data[i][j]=matrix1->data[i][j]-matrix2->data[i][j];
  
  return prod;
}

