mmsearch.c

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

C
2,188
字号
{   int i, j, lat, ku, nr2, nreg=Mmin(nb, maxreg);   double mf;/* * For large number of registers, search only near-square cases */   if (maxreg > 64)   {      searchnu_nu(pre, nb, maxreg, Fku, muladd, pfA, LAT, NO1D, mfB, nbB, muB,                  nuB, kuB, latB);      return;   }   for (i=1; i <= nreg; i++)   {      nr2 = maxreg / i;      if (nr2 > nb) nr2 = nb;      for (j=1; j <= nreg; j++)      {         if ( (((i==1) && (j > 4)) || ((j==1) && (i > 4))) && NO1D) continue;         lat = GetGoodLat(muladd, nb, i, j, ku, LAT);         lat = Mmin(lat, LAT);         if (i*j+i+j+(!muladd)*lat < maxreg)            kucases(pre, muladd, nb, i, j, pfA, LAT, Fku,                    mfB, nbB, muB, nuB, kuB, latB);      }   }}#endifvoid FindKU(char pre, int muladd, int pfA, int LAT, int nb, int mu, int nu,            double *mfB, int *kuB, int *latB)/* * For best case, try various ku's */{   int k, lat, size, linesize;   double mf;   fprintf(stderr, "Confirming K-loop unrollings for chosen NB:\n");   mf = mms_case(pre, muladd, nb, mu, nu, nb, pfA, LAT);   if (mf > *mfB)   {      *kuB = nb;      *mfB = mf;      *latB = LAT;   }/* * Try 2, 4, 6, 8 */   for (k=2; k < 8; k += 2)   {      lat = GetGoodLat(muladd, nb, mu, nu, k, *latB);      mf = mms_case(pre, muladd, nb, mu, nu, k, pfA, lat);      if (mf > *mfB)      {         *latB = lat;         *kuB = k;         *mfB = mf;      }   }/* * Try all unrollings between cache-line size and nb/2, in multiples of * the cacheline size */   if (pre == 's' || pre == 'c')      linesize = ATL_L1LS / sizeof(float);   else      linesize = ATL_L1LS / sizeof(double);   if (linesize < 4) linesize = 4;   for (k=linesize; k < nb; k += linesize)   {      if (k >= nb/2) k = nb;      lat = GetGoodLat(muladd, nb, mu, nu, k, *latB);      mf = mms_case(pre, muladd, nb, mu, nu, k, pfA, lat);      if (mf > *mfB)      {         *latB = lat;         *kuB = k;         *mfB = mf;      }   }}void FindLAT(char pre, int pfA, int maxlat, int nb, int muladd,             int mu, int nu, int ku, double *mfB, int *latB){   int i, lat;   double mf;/* * Right now, search does not do accumulator expansion for small (mu, nu), * so there is no need to search for MAC */   if (muladd) return;   fprintf(stderr, "\nConfirming latency factors for chosen parameters:\n");   for (i=1; i <= maxlat; i++)   {      lat = GetGoodLat(muladd, nb, mu, nu, ku, i);      if (lat == i)      {         mf = mms_case(pre, muladd, nb, mu, nu, ku, pfA, lat);         if (mf > *mfB)         {            *mfB = mf;            *latB = i;         }      }   }   fprintf(stderr, "\n\n   Best latency factor=%d\n\n", *latB);}int ProbeLatency(char pre, int nr, int nb, int muladd, int mu, int nu)/* * Finds a good setting for latency, assuming mu*nu > pipeline.  Uses * the minimum latency that gets good performance.  If latency has no * real affect on performance, returns latency of 1 (to avoid multiple probs * and minimize register waste) */{   int maxlat, i;   double mf0, mf;   if (muladd) return(1);   maxlat = Mmin(nr/2, 16);/* * Try latencies between 1 and 16, stopping anytime performance does not * improve; always unroll loop all way to help avoid having compiler pipeline * and so that all lats can be tried */   fprintf(stderr, "\nProbing for a good latency value:\n");   mf0 = mms_caseIC(pre, 0, nb, mu, nu, nb, 0, 1);   fprintf(stderr, "   lat = %d, mf=%.2f\n", 1, mf0);   for (i=2; i <= maxlat; i++)   {      mf = mms_caseIC(pre, 0, nb, mu, nu, nb, 0, 1);      fprintf(stderr, "   lat = %d, mf=%.2f\n", i, mf);      if (mf < mf0*1.01) /* w/o 1% improvement, latency not worth increasing */         break;      mf0 = mf;   }   i--;   fprintf(stderr, "lat=%d selected!\n", i);   return(i);}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);}static int GuessSmallNB(char pre, int L1Size, int mu, int nu)/* * Returns a small nb useful for in-cache timings */{   int imult, nb, size;   size = (pre == 'd' || pre == 'z') ? ATL_dsize : ATL_ssize;   L1Size /= size;   imult = Mylcm(mu, nu);/* * Try to get a block factor where A, B & C all fit into cache */   for (nb=imult; 3*nb*nb < L1Size; nb += imult);   nb -= imult;/* * If block to small, settle for fitting one block comfortably in cache */   if (nb < 28)   {      for (; nb*nb+(mu+nu)*nb*2 < L1Size; nb += imult);      nb -= imult;   }   fprintf(stderr, "L1Size=%d, pre=%c, Smallnb=%d\n", L1Size, pre, nb);   assert(nb);   return(nb);}void ProbeFPU(char pre, int L1Size, int nreg, int *muladd0, int *lat0)/* * Estimates good muladd and latency for matmul */{   double mf0, mf1;   int i, mu, nu, muladd_r, lat_r, nb, imult, lat;   char upre=pre;   void GetMulAdd(char pre, int *MULADD, int *lat);   if (pre == 'c') upre = 's';   else if (pre == 'z') upre = 'd';/* * Get muladd & latency for register-to-register code */   GetMulAdd(upre, &muladd_r, &lat_r);   FindMUNU(0, lat_r, (nreg > 16) ? nreg-2 : nreg, 0, &mu, &nu);/* * Find good nb to use with these parameters */   nb = GuessSmallNB(pre, L1Size, mu, nu);/* * Compute best latency setting for separate multiply and add */   lat = ProbeLatency(pre, nreg, nb, 0, mu, nu);/* * Get mu,nu and nb to use with real matmul-detected latency */   FindMUNU(0, lat, (nreg > 16) ? nreg-2 : nreg, 0, &mu, &nu);   nb = GuessSmallNB(pre, L1Size, mu, nu);/* * Time separate and combined mul/add. * NOTE: may slightly disadvantage muladd=1 case, as it uses the mu,nu set * with muladd=0 (using lat extra regs), but this gives us the same blocking * factor.  To offset this, require mf0 by 3% better to avoid using muladd=1 */   mf0 = mms_caseIC(pre, 0, nb, mu, nu, nb, 0, lat);   mf1 = mms_caseIC(pre, 1, nb, mu, nu, nb, 0, lat);   if (mf0 >= 1.03*mf1)   {      *muladd0 = 0;      *lat0 = lat;   }   else   {      *muladd0 = 1;      *lat0 = lat_r;   }   fprintf(stderr, "\n\nMATMUL FPU PROBE RESULTS: muladd=%d, lat=%d (%.2f) selected over (%.2f)!!\n",           *muladd0, *lat0,           *muladd0 == 0 ? mf0 : mf1, *muladd0 == 0 ? mf1 : mf0);}int CheckUser(char pre, double adv, double gmf, int gnb, double *umf)/* * Checks if user case is better than generated, and if so, return umb */{   FILE *fp;   char fnam[128];   int i, unb;   double umflop;   sprintf(fnam, "res/%cuMMRES", pre);   if (!FileExists(fnam))  /* need to run user search */   {      sprintf(fnam, "make RunUMMSearch pre=%c nb=%d\n", pre, gnb);      assert(system(fnam) == 0);      sprintf(fnam, "res/%cuMMRES", pre);   }   fp = fopen(fnam, "r");   assert(fp);   assert(fgets(fnam, 128, fp));   assert(fgets(fnam, 128, fp));   fclose(fp);   sscanf(fnam, " %d %d %lf", &i, &unb, &umflop);   if (i >= 0 && umflop < 0.0)  /* need to retime */   {      sprintf(fnam, "make RunUMMSearch pre=%c nb=0\n", pre);      assert(system(fnam) == 0);      sprintf(fnam, "res/%cuMMRES", pre);      fp = fopen(fnam, "r");      assert(fp);      assert(fgets(fnam, 128, fp));      assert(fgets(fnam, 128, fp));      fclose(fp);      sscanf(fnam, " %d %d %lf", &i, &unb, &umflop);   }   fprintf(stdout, "\nBEST USER CASE: NB=%d, MFLOP=%.2f\n", unb, umflop);   if (umf) *umf = umflop;   if (adv*gmf > umflop || i < 1) return(gnb);   else return(unb);}int GetNO1D(char pre, int nreg, int nb, int MULADD, int pfA, int LAT){   int lat, NO1D=0;   double mf0, mf1, mf;/* * Always do 1-D cases for 2-op assembler! */   #ifdef TWO_OP_ASM      return(0);   #endif   if (pre == 'z') pre = 'd';   else if (pre == 'c') pre = 's';   lat = GetGoodLat(MULADD, nb, 3, 3, 1, LAT);   if (nreg >= 15+(!MULADD)*Mmax(LAT,lat))   {      mf0 = mms_case(pre, MULADD, nb, 3, 3, 1, pfA, lat);      mf1 = mms_case(pre, MULADD, nb, 3, 3, nb, pfA, LAT);      mf = Mmax(mf1, mf0);      mf0 = mms_case(pre, MULADD, nb, 9, 1, 1, pfA, lat);      if (mf0 > mf) NO1D = 0;      else if (mms_case(pre, MULADD, nb, 9, 1, nb, pfA, LAT) > mf) NO1D = 0;      else if (mms_case(pre, MULADD, nb, 1, 9, nb, pfA, LAT) > mf) NO1D = 0;      else if (mms_case(pre, MULADD, nb, 1, 9, 1, pfA, lat) > mf) NO1D = 0;      else NO1D = 1;   }   return(NO1D);}double SearchNBs(char pre,      /* s, d */                 int MA,        /* muladd */                 int Fku,/* =0, try both ku=1 and ku=KB, else try only ku=Fku */                 int nNBs,      /* # of NBs in NB array */                 int *NBs,      /* array of NBs to search */                 int mu,        /* M-loop unrolling to use */                 int nu,        /* N-loop unrolling to use */                 int pfA,       /* 0: no prefetch of A */                 int lat,       /* approx latency to enforce */                 int *latBo,    /* latency used by best kernel */                 int *kuBo)     /* KU used by best kernel *//* * RETURNS: Mflop of case that got the best performance (best NB is returned *          in first position of NBs array). */{   double mf, mfB;   int i, j, k, tlat, nb, ku, kuB, latB;   if (Fku < 0) Fku = 0;   tlat = lat; kuB = Fku; latB = lat;   i = 0; mfB = 0.0;   for (k=0; k != nNBs; k++)   {      nb = NBs[k];      ku = Fku ? Fku : nb;      mf = mms_case(pre, MA, nb, mu, nu, ku, pfA, lat);      if (mf > mfB)      {         mfB = mf;         kuB = ku;         latB = lat;         i = k;      }      if (Fku == 0)  /* try no K-loop unrolling */      {         tlat = GetGoodLat(MA, nb, mu, nu, 1, lat);         mf = mms_case(pre, MA, nb, mu, nu, 1, pfA, tlat);         if (mf > mfB)         {            mfB = mf;            kuB = 1;            latB = tlat;            i = k;         }      }   }/* * Put best-performing NB first in NB array */   if (i)   {      j = NBs[i];      NBs[i] = NBs[0];      NBs[0] = j;   }   fprintf(stderr, "NB=%d selected:\n", NBs[0]);   *kuBo = kuB;   *latBo = latB;   return(mfB);}void gmmsearch(char pre, int MULADD, int Fku, int nNBs, int *NBs, int nreg,               int LAT, int Fnb)/* * Does real generated mmsearch */{   int latB, muB, nuB, kuB, nbB;   int i, j, k, NB, ku, nb, lat=LAT, nNB=nNBs, NO1D=0;   int FFetch, ifetch, nfetch, muladd, pfA;   double mf, mfB, mf1;   char ln[32];   FILE *fp;   sprintf(ln, "res/%cgMMRES", pre);   if (FileExists(ln)) /* already have needed result */   {      GetInstLogFile(ln, pre, &muladd, &pfA, &lat, &nb, &muB, &nuB, &kuB,                     &FFetch, &ifetch, &nfetch, &mf);

⌨️ 快捷键说明

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