/***********************************************************************:

 SuperCon 2001 本選回答 (nothing版)    [2001/08/02] by tnj @ nothing
 
 ----------------------------------------------------------------------
 
 なんか答え合わないらしいんすけど、コレだとVerify通ります。
 問題BでもVerify通ってるんです。何で問題Aは通らないんですか。
 
 計算部分は弄ってないので、合わないはずはないんですが…。
 おかしい。くそっ。
 
 round関数弄ったと言ったんですが、弄ってませんでした。
 だから、roundが原因とかそんなことはないはずです。

 だいたいどのくらい誤差出てるんですか。
 小数点以下6桁は合ってるハズなんですよ。こっちでVerify通ってるから。
 どのくらい誤差でてるかわかんないんですか、って訊いたんですが。
 結局答えてもらえませんでした。
 
 というか、問題Cの計算時間出してくださいよ。非道いですよ。
 …なんて。
 
 どちらにしろこのスピードじゃ5位以下でしたが。
 負け惜しみです。(;´Д`)
 
 なお、MPI関連のコードを削るだけで単体で動作するはずです。
 あと、numprocs = 1; myprocid = 0; とでもすれば。
 
:***********************************************************************/

#include <stdio.h>
#include <math.h>
#include <stdlib.h>
#include <time.h>
#include <assert.h>

#include "mpi.h"

#define root 0

double DT;
int n_steps;
int n_particles;

double *im,*ix,*iy,*ivx,*ivy,*iax,*iay;
double *tax,*tay;
double cutoff_distance_2 = 1024.0;

void   supercon2k1_get_config(int problem_num, int * size_p, int * n_steps_p, double * dt_p);
void   supercon2k1_get_data(int problem_num, double * a);
void   supercon2k1_verify(int problem_num, double * a);

void   do_step();

int    numprocs,myprocid;
int    mypart_start,mypart_end;
int    mypart_count;

int    partstart[32],partnum[32];

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

// #define round(double x) (x<0.0) ? ceil(x+0.5) : floor(x+0.5)


double round(double x) {
  if (x < 0.0) {
    return -floor (-x + .5);
  } else {
    return floor (x + .5);
  }
} 


int main(int argc, char ** argv)
{
  int problem_num, i, j, stnum;
  double *a,*ap;
  int allocsize;

  MPI_Init(&argc,&argv);
  MPI_Comm_size(MPI_COMM_WORLD,&numprocs);
  MPI_Comm_rank(MPI_COMM_WORLD,&myprocid);

  if (myprocid == root) {

    //printf("root launched!\nnumprocs=%d\n",numprocs);

    if(argc < 2){
      fprintf(stderr,"please give me  question number.\n");
      exit(1);
    }
    
    problem_num = atoi(argv[1]);
    supercon2k1_get_config(problem_num,&n_particles,&n_steps,&DT);
    printf("problem_num=%d: n_particle=%d, n_steps=%d, DT=%g\n", problem_num, n_particles, n_steps, DT);
  }
  
  MPI_Bcast(&n_particles, 1, MPI_INT, 0, MPI_COMM_WORLD);
  MPI_Bcast(&n_steps, 1, MPI_INT, 0, MPI_COMM_WORLD);
  MPI_Bcast(&DT, 1, MPI_DOUBLE, 0, MPI_COMM_WORLD);
  
// memory allocation
  
  allocsize = sizeof(double) * n_particles;
  a = (double *)malloc(5 * allocsize);         /* linear */
  
  im = (double *)malloc(allocsize);
  ix = (double *)malloc(allocsize);
  iy = (double *)malloc(allocsize);
  ivx = (double *)malloc(allocsize);
  ivy = (double *)malloc(allocsize);
  iax = (double *)malloc(allocsize);
  iay = (double *)malloc(allocsize);
  
    tax = (double *)malloc(allocsize);
    tay = (double *)malloc(allocsize);

  if(a == NULL || iay == NULL){
    fprintf(stderr,"proc%d: can't allocate memory...\n",myprocid);
    exit(1);
  }
  
  if (myprocid == root) {

    //printf("proc%d: getting data...\n",myprocid);

    supercon2k1_get_data(problem_num,a);
    ap = a;
    for(i = 0; i < n_particles; i++) {
      im[i] = *ap++;
      ix[i] = *ap++;
      iy[i] = *ap++;
      ivx[i] = *ap++;
      ivy[i] = *ap++;
    }

    /* root allocates jobs */
    j=0;
    stnum = numprocs*(numprocs+1)/2;
    for(i=0;i<numprocs;i++){
      partstart[i]=j;
      j+=n_particles*(i+1)/stnum+1;
      if (j>n_particles) j=n_particles;
      partnum[i]=j-partstart[i];
    }

    /*


        j=0;
    stnum = (int)((n_particles) / numprocs);
    for(i=0;i<numprocs;i++){
      partstart[i]=j;
      j+=stnum;
      if (j>n_particles) j=n_particles;
      partnum[i]=j-partstart[i];

    }
      */
  }

  /* broadcast part allocation arrays */
  MPI_Bcast((void *)partstart, numprocs, MPI_INT, root, MPI_COMM_WORLD);
  MPI_Bcast((void *)partnum, numprocs, MPI_INT, root, MPI_COMM_WORLD);

  mypart_start = partstart[myprocid];
  mypart_count = partnum[myprocid];
  mypart_end   = mypart_start + mypart_count - 1;
  
  //printf("proc%d: caliculate %d - %d [%d]\n",myprocid,mypart_start,mypart_end,mypart_count);

  /* broadcast arrays */

  MPI_Bcast(im, n_particles, MPI_DOUBLE, root, MPI_COMM_WORLD);
  MPI_Bcast(ix, n_particles, MPI_DOUBLE, root, MPI_COMM_WORLD);
  MPI_Bcast(iy, n_particles, MPI_DOUBLE, root, MPI_COMM_WORLD);
  MPI_Bcast(ivx, n_particles, MPI_DOUBLE, root, MPI_COMM_WORLD);
  MPI_Bcast(ivy, n_particles, MPI_DOUBLE, root, MPI_COMM_WORLD);

  if (myprocid==root) printf("proc0: array broadcasted.\n");
        
  for(i = 0; i < n_steps; i++){
    do_step();

//    fprintf(stdout,"!%d\n",i);
//    if ((i+1) % 10 == 0) fprintf(stdout," ");
//    if ((i+1) % 50 == 0) fprintf(stdout,"[%5d]\n",i+1);
//    fflush(stdout);

  
  }


  //printf("proc%d: complete.\n",myprocid);

  MPI_Gatherv(ivx+mypart_start,mypart_count,MPI_DOUBLE,ivx,partnum,partstart,MPI_DOUBLE,root,MPI_COMM_WORLD);
  MPI_Gatherv(ivy+mypart_start,mypart_count,MPI_DOUBLE,ivy,partnum,partstart,MPI_DOUBLE,root,MPI_COMM_WORLD);

  if (myprocid == root) {

    //printf("proc0: prepering verify...\n");

  // convert to linear array (for verify)
    ap = a;
    for(i = 0; i < n_particles; i++) {
      *ap++ = im[i];
      *ap++ = ix[i];
      *ap++ = iy[i];
      *ap++ = ivx[i];
      *ap++ = ivy[i];
    }
  
  // verify answer
    supercon2k1_verify(problem_num, (double *)a);
  
  }

  MPI_Finalize();
  exit(0);
}

void do_step()
{
  int i, j;
  double dx, dy;
  double norm2;
  double acc;

//  MPI_Barrier(MPI_COMM_WORLD);

  for(i = n_particles; i >= 0; i--) {
    iax[i]=0.0; iay[i]=0.0; 
  }

  if (myprocid == root) {
    for(i = n_particles; i >= 0; i--) {
      tax[i]=0;tay[i]=0;
    }
  }
	
  for(i = mypart_end; i >= mypart_start; i--) {
    for(j = n_particles ; j > i; j--){
//      if(i == j) continue;

      dx = ix[j] - ix[i];
      dy = iy[j] - iy[i];

      if(dx*dx + dy*dy < cutoff_distance_2) {
	if(dx>0) dx = ceil(dx); else dx = floor(dx);
	if(dy>0) dy = ceil(dy); else dy = floor(dy);
	
	norm2 = dx*dx + dy*dy;
	assert(norm2);
	
	if(norm2 < cutoff_distance_2) {
	  acc = im[j] / norm2;
	  iax[i] += round(acc * (dx / sqrt(norm2)));
	  iay[i] += round(acc * (dy / sqrt(norm2)));
	  acc = im[i] / norm2;
	  iax[j] -= round(acc * (dx / sqrt(norm2)));
	  iay[j] -= round(acc * (dy / sqrt(norm2)));
	} 
      }
    }
  }

  //printf("// %5d\n",myprocid);

  //MPI_Barrier(MPI_COMM_WORLD);

  //printf("proc%d:trap0\n",myprocid);

  MPI_Reduce(iax,tax,n_particles,MPI_DOUBLE,MPI_SUM,0,MPI_COMM_WORLD);
  MPI_Reduce(iay,tay,n_particles,MPI_DOUBLE,MPI_SUM,0,MPI_COMM_WORLD);

  MPI_Bcast(tax,n_particles,MPI_DOUBLE,0,MPI_COMM_WORLD);
  MPI_Bcast(tay,n_particles,MPI_DOUBLE,0,MPI_COMM_WORLD);

  for(i = mypart_end; i >= mypart_start; i--){
    ix[i] += ivx[i] * DT;
    iy[i] += ivy[i] * DT;
    ivx[i] += tax[i] * DT;
    ivy[i] += tay[i] * DT;
  }

  /* writeback position arrays of my part */

  //MPI_Barrier(MPI_COMM_WORLD);
  
  //printf("proc%d:trap2\n",myprocid);

  MPI_Allgatherv(ix+mypart_start,mypart_count,MPI_DOUBLE,ix,partnum,partstart,MPI_DOUBLE,MPI_COMM_WORLD);
  MPI_Allgatherv(iy+mypart_start,mypart_count,MPI_DOUBLE,iy,partnum,partstart,MPI_DOUBLE,MPI_COMM_WORLD);

  //MPI_Barrier(MPI_COMM_WORLD);


    #ifdef SC2001OUT
      for (i=0; i<number_of_stars; i++)   /* don't worry the order of */
          printf("%d %d %5f %5f\n", rank, step, x, y);  /* the output */
    #endif



}
