ummsearch.c

来自「基于Blas CLapck的.用过的人知道是干啥的」· C语言 代码 · 共 1,470 行 · 第 1/3 页

C
1,470
字号
 * Reduces table to best in each catagory */{   MULTHEAD *mh, *mhnext;   const double advant=0.03;   for (mh=imhead; mh; mh = mhnext)   {      mhnext = mh->next;      ReduceRouts(pre, which, mh, advant, nb);   }   ReduceMults(pre, which, advant, nb);}int ummtstcase0(   char pre,                  /* type prefix */   int M, int N, int K,       /* problem sizes to test */   int mb, int nb, int kb,    /* 0: variable NB, else fixed cpp macro of NB */   int lda, int ldb, int ldc, /* leading dims */   int muladd, int lat,       /* muladd and latency settings */   int mu, int nu, int ku,    /* unrolling factors */   char *fnam,                /* file name to compile */   char *MCC, char *MMFLAGS   /* NULL : use defaults, else comp to use */){   char ln[512];   int i;   char ch;   if (pre == 'c' || pre == 'z')      i = sprintf(ln, "make cmmutstcase mmrout=CASES/%s csC=2 ", fnam);   else i = sprintf(ln, "make mmutstcase mmrout=CASES/%s ", fnam);   if (MCC)   {      ch = (pre == 's' || pre == 'c') ? 'S' : 'D';      i += sprintf(ln+i, "%cMC=\"%s\" %cMCFLAGS=\"%s\" ", ch, MCC, ch, MMFLAGS);   }   i += sprintf(ln+i, "pre=%c muladd=%d lat=%d M=%d N=%d K=%d mb=%d nb=%d kb=%d mu=%d nu=%d ku=%d lda=%d ldb=%d ldc=%d ",                pre, muladd, lat, M, N, K, mb, nb, kb, mu, nu, ku,                lda, ldb, ldc);   i += sprintf(ln+i, "\n");   fprintf(stdout, "%s", ln);   return(system(ln) == 0);}static int Mylcm(const int M, const int N)/* * Returns least common multiple (LCM) of two positive integers M & N by * computing greatest common divisor (GCD) and using the property that * M*N = GCD*LCM. */{   register int tmp, max, min, gcd=0;   if (M != N)   {      if (M > N) { max = M; min = N; }      else { max = N; min = M; }      if (min > 0)  /* undefined for negative numbers */      {         do  /* while (min) */         {            if ( !(min & 1) ) /* min is even */            {               if ( !(max & 1) ) /* max is also even */               {                  do                  {                     min >>= 1;                     max >>= 1;                     gcd++;                     if (min & 1) goto MinIsOdd;                  }                  while ( !(max & 1) );               }               do min >>=1 ; while ( !(min & 1) );            }/* *          Once min is odd, halve max until it too is odd.  Then, use *          property that gcd(max, min) = gcd(max, (max-min)/2) *          for odd max & min */MinIsOdd:            if (min != 1)            {               do  /* while (max >= min */               {                  max -= (max & 1) ? min : 0;                  max >>= 1;               }               while (max >= min);            }            else return( (M*N) / (1<<gcd) );            tmp = max;            max = min;            min = tmp;         }         while(tmp);      }      return( (M*N) / (max<<gcd) );   }   else return(M);}int GetUserNB(char pre, int NB, int mb, int nb, int kb)/* * Given ATLAS's preference for NB, finds a good nb based on user's input */{   int i, nelt, mult;   if (mb < 0)   {      if ( (nb < 0 && nb != mb) || (kb < 0 && kb != mb) ) return(0);      else return(-mb);   }   else if (nb < 0)   {      if (kb < 0 && kb != nb) return(0);      else return(-nb);   }   else if (kb < 0) return(-kb);   if (!nb) nb = 1;   if (!mb) mb = 1;   if (!kb) kb = 1;   mult = Mylcm(mb,nb);   mult = Mylcm(kb,mult);   if ( (NB%mb == 0) && (NB%nb == 0) && (NB%kb == 0) ) return(NB);   nelt = L1Elts(pre, 128*1024);   i = ((NB+mult-1)/mult)*mult;   if (i*i < nelt) return(i);   else return((NB/mult)*mult);   return(0);}char *GetUserOutFile(char pre, int ifile, int mb, int nb, int kb){   static char ln[32];   sprintf(ln, "res/%cuser%03d_%dx%dx%d", pre, ifile+1, mb, nb, kb);   return(ln);}double ummcase(char pre, int ifile, int NB){   char outnam[256], fnam[256];   char *MCC, *MMFLAGS;   int iflag, mb, nb, kb, muladd, lat, mu, nu, ku;   assert(GetUserCase(pre, ifile, &iflag, &mb, &nb, &kb, &muladd, &lat,                      &mu, &nu, &ku, fnam, outnam, &MCC, &MMFLAGS));   if (ATL_MMCleanOnly(iflag)) return(0.0); /* don't run if for cleanup only */   return(ummcase0(pre, NB, NB, NB, NB, NB, NB, NB, NB, 0, muladd, lat,                   mu, nu, ku, fnam, MCC, MMFLAGS,                   GetUserOutFile(pre, ifile, NB, NB, NB)));}int utstmmcase(char pre, int ifile, int NB){   char outnam[256], fnam[256];   char *MCC, *MMFLAGS;   int iflag, mb, nb, kb, muladd, lat, mu, nu, ku;   assert(GetUserCase(pre, ifile, &iflag, &mb, &nb, &kb, &muladd, &lat,                      &mu, &nu, &ku, fnam, outnam, &MCC, &MMFLAGS));   return(ummtstcase0(pre, NB, NB, NB, NB, NB, NB, NB, NB, 0, muladd, lat,                      mu, nu, ku, fnam, MCC, MMFLAGS));}int FindBestUser(char pre, int nb0)/* * returns index in <pre>cases.dsc of best user-supplied GEMM, using * a blocking factor as close to nb0 as possible */{   char *MCC, *MMFLAGS;   char ln[256], fnam[256];   double mf, mfbest=0.0;   int ibest=(-1), i, ncases;   int ID, iflag, NB, mb, nb, kb, ma, lat, mu, nu, ku;   ncases = NumUserCases(pre);   for (i=0; i < ncases; i++)   {      ID = GetUserCase(pre, -i, &iflag, &mb, &nb, &kb, &ma, &lat,                       &mu, &nu, &ku, fnam, ln, &MCC, &MMFLAGS);      assert(ID > 0);      NB = GetUserNB(pre, nb0, mb, nb, kb);      if (NB)      {         mf = ummcase(pre, ID, NB);         if (mf > mfbest)         {            if (utstmmcase(pre, ID, NB))            { /* test kernel before accepting */               ibest = ID;               mfbest = mf;            }         }         fprintf(stdout, "%3d. NB=%3d, rout=%40s, MFLOP=%.2f\n",                 i, NB, fnam, mf);      }   }   return(ibest);}void PrintUsage(char *fnam){   fprintf(stderr, "\nUSAGE: %s -p <pre> -n <nb>\n\n", fnam);   exit(-1);}void GetFlags(int nargs, char **args, char *pre, int *nb0,              enum CLEAN_WHICH *which){   int i;   char ch;   *nb0 = 0;   *pre = 'd';   *which = CleanNot;   for (i=1; i < nargs; i++)   {      if (args[i][0] != '-') PrintUsage(args[0]);      switch(args[i][1])      {      case 'p':         i++;         ch = tolower(args[i][0]);         if (ch == 'd' || ch == 's' || ch == 'z' || ch == 'c') *pre = ch;         else PrintUsage(args[0]);         break;      case 'n':         *nb0 = atoi(args[++i]);         break;      case 'C':         switch(args[++i][0])         {         case 'm':            *which = CleanM;            break;         case 'n':            *which = CleanN;            break;         case 'k':            *which = CleanK;            break;         default:            *which = CleanNot;         }         break;      default:         PrintUsage(args[0]);      }   }}int FindBestNB(pre, icase){   double mf, mfb;   int iflag, mb, nb, kb, muladd, lat, mu, nu, ku;   int i, j, mult, iret=0;   char fnam[ROUTLEN], *MCC, *MMFLAGS;   assert(GetUserCase(pre, icase, &iflag, &mb, &nb, &kb, &muladd, &lat,                      &mu, &nu, &ku, fnam, fnam, &MCC, &MMFLAGS));   if (mb < 0) iret = -mb;   else if (nb < 0) iret = -nb;   else if (kb < 0) iret = -kb;   else   {/* *    Find mult necessary to satisfy user constraints on mb, nb, and kb */      if (mb != 0) i = mb;      else i = 1;      if (nb != 0) j = nb;      else j = 1;      mult = Mylcm(i, j);      if (kb != 0) i = kb;      else i = 1;      mult = Mylcm(i, mult);/* *    Make sure mult is a multiple of Cachelen as well */      if (pre == 's' || pre == 'c') i = ATL_Cachelen / ATL_ssize;      else i = ATL_Cachelen / ATL_dsize;      if (i > 1) mult = Mylcm(i, mult);      fprintf(stdout, "\nFINDING BEST BLOCKING FACTOR FOR CASE %d, MUL=%d:\n",              icase, mult);      j = L1Elts(pre, 128);      for (i=mult; i < 16; i += mult);      mfb = 0.0;      iret = i;      for (; ((i*i <= j) && (i <= MAX_NB)); i += mult)      {         mf = ummcase(pre, icase, i);         if (mf > mfb)         {            mfb = mf;            iret = i;         }         fprintf(stdout, "   NB=%d: %.2f MFLOP\n", i, mf);      }   }   fprintf(stdout, "BEST BLOCKING FACTOR FOR CASE %d:  %d\n", icase, iret);   return(iret);}int NumMults(void){   MULTHEAD *mh;   int i;   for(i=0, mh=imhead; mh; mh = mh->next, i++) if (mh->rn->next) i++;   return(i);}void Printpline(char pre, enum CLEAN_WHICH which, int nb, int imult,                ROUTNODE *rn, FILE *fp){   int NB[3];   GetPNB(pre, which, rn->icase, nb, imult, NB);   fprintf(fp, "%4d  %5d  %5d %3d  %3d  %3d  %3d  %8.2f  %s\n",           imult, rn->icase, rn->fixed, nb, NB[0], NB[1], NB[2],           rn->mflop, rn->rout);}void CreatepUMMOut(char pre, enum CLEAN_WHICH which, int nb){   MULTHEAD *mh;   char cwh[3] = {'M', 'N', 'K'};   char fnam[128];   int NB[3];   FILE *fp;   sprintf(fnam, "res/%cuClean%c", pre, cwh[which]);   fp = fopen(fnam, "w");   assert(fp);   fprintf(fp, "MULT  ICASE  FIXED  NB  NB0  NB1  NB2     MFLOP  ROUT\n");   fprintf(fp, "%d\n", NumMults());   for (mh=imhead; mh; mh = mh->next)   {      Printpline(pre, which, nb, mh->imult, mh->rn, fp);      if (mh->rn->next) Printpline(pre, which, nb, mh->imult, mh->rn->next, fp);   }}void FindUClean(char pre, int nb0, enum CLEAN_WHICH which){   BuildTable(pre, which, nb0);   PrintTable(stdout);   TimeTable(pre, which, nb0);   PrintTable(stdout);   ReduceTable(pre, which, nb0);   PrintTable(stdout);   CreatepUMMOut(pre, which, nb0);   KillAllMultNodes();}void CreateUMMOut(char pre, int icase, int NB, double mf){   char fnam[ROUTLEN], auth[AUTHLEN], *MCC, *MMFLAGS;   int iflag, mb, nb, kb, muladd, lat, mu, nu, ku;   FILE *fp;   sprintf(auth, "res/%cuMMRES", pre);   fp = fopen(auth, "w");   assert(fp);   if (icase > 0)   {      assert(GetUserCase(pre, icase, &iflag, &mb, &nb, &kb, &muladd, &lat,                         &mu, &nu, &ku, fnam, auth, &MCC, &MMFLAGS));   }   else   {      mf = -1.0;      strcpy(auth, "Nobody");      strcpy(fnam, "Nocomp");   }   fprintf(fp, "CASE  NB    MFLOP  ROUTINE\n");   fprintf(fp, "%4d %3d %8.2f  \"%.64s\" \"%.64s\"\n",           icase, NB, mf, fnam, auth);   fclose(fp);}void FindUMM(char pre, int nb0){   double mf, mf0;   int nb, icase;   icase = FindBestUser(pre, nb0);   if (icase >= 0)   {      mf0 = GetRes(GetUserOutFile(pre, icase, nb0, nb0, nb0));      nb = FindBestNB(pre, icase);      mf = GetRes(GetUserOutFile(pre, icase, nb, nb, nb));      if (mf <= mf0)      {         nb = nb0;         mf = mf0;      }   }   else   {      nb = nb0;      mf = -1.0;   }   CreateUMMOut(pre, icase, nb, mf);   fprintf(stdout, "\nBEST USER CASE %d, NB=%d: %.2f MFLOP\n\n", icase, nb, mf);}void RunTimes(char pre){   double mf;   FILE *fp;   int j, icase, nb;   char ln[128];   sprintf(ln, "res/%cuMMRES", pre);   if (FileExists(ln))   {      fp = fopen(ln, "r");      assert( fgets(ln, 128, fp) != NULL );      assert( fscanf(fp, " %d %d", &icase, &nb) == 2);      fclose(fp);      if (icase >= 0)      {         mf = ummcase(pre, icase, nb);         CreateUMMOut(pre, icase, nb, mf);         fprintf(stdout, "\nBEST USER CASE %d, NB=%d: %.2f MFLOP\n\n",                 icase, nb, mf);      }   }   else   {      sprintf(ln, "res/%cMMRES", pre);      if (!FileExists(ln)) sprintf(ln, "res/%cgMMRES", pre);      GetInstLogFile(ln, pre, &j, &j, &j, &nb, &j, &j, &j, &j, &j, &j, &mf);      FindUMM(pre, nb);   }}void GetuMMRES(char pre){   double mf;   FILE *fp;   int j, icase, nb;   char ln[128];   sprintf(ln, "res/%cuMMRES", pre);   if (FileExists(ln))   {      fp = fopen(ln, "r");      assert( fgets(ln, 128, fp) != NULL );      assert( fscanf(fp, " %d %d %lf", &icase, &nb, &mf) == 3);      fclose(fp);      if (icase >= 0 && mf <= 0.0)      {         mf = ummcase(pre, icase, nb);         CreateUMMOut(pre, icase, nb, mf);         fprintf(stdout, "\nBEST USER CASE %d, NB=%d: %.2f MFLOP\n\n",                 icase, nb, mf);      }   }   else   {      sprintf(ln, "res/%cMMRES", pre);      if (!FileExists(ln)) sprintf(ln, "res/%cgMMRES", pre);      GetInstLogFile(ln, pre, &j, &j, &j, &nb, &j, &j, &j, &j, &j, &j, &mf);      FindUMM(pre, nb);   }}main(int nargs, char **args){   int i, nb0;   enum CLEAN_WHICH which;   char pre;   GetFlags(nargs, args, &pre, &nb0, &which);   if (nb0 <= 0) GetuMMRES(pre);   else if (which == CleanNot) FindUMM(pre, nb0);   else FindUClean(pre, nb0, which);   exit(0);}

⌨️ 快捷键说明

复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?