mmsearch.c

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

C
2,188
字号
      if (mf <= 0.0)      {         mf = mmcase(NULL, pre, "JIK", 'T', 'N', nb, nb, nb, nb, nb, nb,                     nb, nb, 0, muB, nuB, kuB, muladd, pfA, lat, 1, 1, 1, 2,                     FFetch, ifetch, nfetch);         PutInstLogFile1(ln, pre, muladd, pfA, lat, nb, muB, nuB, kuB, FFetch,                         ifetch, nfetch, mf);      }      return;   }/* * Try not to tempt fate by using all registers */   if (nreg > 16) i = nreg-2;   else i = nreg;   FindMUNU(MULADD, lat, i, 0, &muB, &nuB);/* * First, find a good NB, with no prefetch */   pfA = 0;   mfB = SearchNBs(pre, MULADD,  Fku, nNBs, NBs, muB, nuB, 0,                   lat, &latB, &kuB);   nbB = NBs[0];/* * Now, find if prefetching helps this kernel; may want to user different * NB if prefetch is a win (esp., smaller NB) */   mf  = SearchNBs(pre, MULADD,  Fku, nNBs, NBs, muB, nuB, 1,                   lat, &lat, &ku);   nbB = NBs[0];   if (mf > mfB)   {      fprintf(stderr, "\nPrefetch kernel %.2f faster.\n", mf/mfB);      mfB = mf;      latB = lat;      kuB = ku;      pfA = 1;   }   else fprintf(stderr, "\nNon-prefetch kernel %.2f faster.\n", mfB/mf);   if (!Fnb) nNB = 1;   if (MULADD)      fprintf(stderr, "\nCombined multiply add, latency factor=%d, NB=%d ku=%d, chosen; initial MFLOP=%f.  Beginning unroll search:\n", latB, NBs[0], kuB, mfB);   else      fprintf(stderr, "\nSeparate multiply and add, latency factor=%d, NB=%d ku=%d, chosen; initial MFLOP=%f.  Beginning unroll search:\n", latB, NBs[0], kuB, mfB);   NO1D = GetNO1D(pre, nreg, NBs[0], MULADD, pfA, LAT);   if (NO1D) fprintf(stderr, "\n\nSkipping most 1D cases\n\n");   else fprintf(stderr, "\n\nTiming 1D cases\n\n");   for (k=0; k != nNB; k++)   {      NB = NBs[k];      searchmu_nu(pre, NB, nreg, Fku, MULADD, pfA, LAT, NO1D,                  &mfB, &nbB, &muB, &nuB, &kuB, &latB);   }   fprintf(stderr, "\n\nBest case so far: nb=%d, mu=%d, nu=%d, ku=%d, lat=%d; MFLOPS=%.2f.\n",           nbB, muB, nuB, kuB, latB, mfB);   fprintf(stderr, "Trying various other NB and KU settings:\n\n");/* * If we haven't checked all permutations, try other blocking factors */   nb = nbB;   if (!Fnb)   {      if (nNBs > 1) fprintf(stderr, "Trying various blocking factors:\n");      mf = mms_case(pre, MULADD, NBs[0], muB, nuB, kuB, pfA, latB);      for (k=0; k < nNBs; k++)      {         NB = NBs[k];         if (Fku == -1) ku = NB;         else if (Fku) ku = Fku;         else if (kuB == nbB) ku = NB;         else ku = kuB;         if (ku != NB) lat = GetGoodLat(MULADD, NB, muB, nuB, ku, latB);         else lat = latB;         mf = mms_case(pre, MULADD, NB, muB, nuB, ku, pfA, lat);         if (mf > mfB)         {            kuB = ku;            mfB = mf;            nbB = NB;            latB = lat;         }      }   }   if (nb != nbB) fprintf(stderr, "\nNew block factor of %d chosen!!\n\n", nbB);   NB = nbB;/* * Try all ku's, and then valid latencies */   FindKU(pre, MULADD, pfA, LAT, nbB, muB, nuB, &mfB, &kuB, &latB);   FindLAT(pre, pfA, MAXLAT, nbB, MULADD, muB, nuB, kuB, &mfB, &latB);/* * Make sure MULADD is correct */   lat = GetGoodLat(!MULADD, nbB, muB, nuB, kuB, latB);   mf = mms_case(pre, !MULADD, nbB, muB, nuB, kuB, pfA, lat);   if (mf > mfB*1.02)   {      fprintf(stderr, "\n\nMULADD MAY BE WRONG!!, old=%f, new=%f\n", mfB, mf);      MULADD = !MULADD;   }/* * See if swapping prefetch helps now */   mf = mms_case(pre, MULADD, nbB, muB, nuB, kuB, !pfA, lat);   if (mf > mfB*1.01)   {      fprintf(stderr, "\n\nPREFETCH SWAPPED TO %d\n\n", pfA);      pfA = !pfA;      mfB = mf;   }/* * Try various fetch patterns */   FindFetch('T', 'N', pre, nbB, nbB, nbB, muB, nuB, kuB, MULADD, pfA, latB,             &FFetch, &ifetch, &nfetch);   fprintf(stdout,   "BEST GENERATED CASE: nb=%d, ma=%d, lat=%d mu=%d, nu=%d, ku=%d -- %.2f\n",           nbB, MULADD, latB, muB, nuB, kuB, mfB);   sprintf(ln, "res/%cgMMRES", pre);   PutInstLogFile1(ln, pre, MULADD, pfA, latB, nbB, muB, nuB, kuB,                   FFetch, ifetch, nfetch, mfB);}void mmsearch(char pre, int MULADD, int Fku, int nNBs, int *NBs, int nreg,              int LAT, int Fnb){   int latB, muB, nuB, kuB, nbB;   int muladd, nb, ifetch, nfetch, FFetch;   int i, j, k, NB, pfA;   int NO1D;   int umb, unb, ukb, ma;   double mfB, gmf;   char fnam[128];   FILE *fp;   sprintf(fnam, "res/%cMMRES", pre);   if (FileExists(fnam)) /* already have result */   {      GetInstLogFile(fnam, pre, &muladd, &pfA, &latB, &nb, &muB, &nuB, &kuB,                     &FFetch, &ifetch, &nfetch, &mfB);      if (mfB <= 0.0)      {         mfB = mmcase(NULL, pre, "JIK", 'T', 'N', nb, nb, nb, nb, nb, nb,                      nb, nb, 0, muB, nuB, kuB, muladd, pfA, latB, 1, 1, 1, 2,                      FFetch, ifetch, nfetch);         gmmsearch(pre, muladd, Fku, nNBs, NBs, nreg, latB, Fnb);         nbB = nb = CheckUser(pre, 1.02, mfB, nb, NULL);         CreateFinalSumm(pre, muladd, pfA, latB, nb, muB, nuB, kuB,                         FFetch, ifetch, nfetch, mfB);      }      else nbB = CheckUser(pre, 1.02, mfB, nb, NULL);      sprintf(fnam, "res/%cNB", pre);      fp = fopen(fnam, "w");      fprintf(fp, "%d\n%d\n", 1, nbB);      fclose(fp);      return;   }   gmmsearch(pre, MULADD, Fku, nNBs, NBs, nreg, LAT, Fnb);   sprintf(fnam, "res/%cgMMRES", pre);   GetInstLogFile(fnam, pre, &muladd, &pfA, &latB, &nbB, &muB, &nuB, &kuB,                  &FFetch, &ifetch, &nfetch, &mfB);   gmf = mfB;   nb = CheckUser(pre, 1.02, mfB, nbB, NULL);   if (nb != nbB)   {      if (kuB == nbB) kuB = nb;      nbB = nb;      if (nb % muB || nb % nuB)      {         NO1D = GetNO1D(pre, nreg, nb, MULADD, pfA, LAT);         searchmu_nu(pre, nb, nreg, Fku, MULADD, pfA, LAT, NO1D,                     &mfB, &nb, &muB, &nuB, &kuB, &latB);      }      FindKU(pre, MULADD, pfA, LAT, nbB, muB, nuB, &mfB, &kuB, &latB);      FindLAT(pre, pfA, MAXLAT, nbB, MULADD, muB, nuB, kuB, &mfB, &latB);      FindFetch('T', 'N', pre, nbB, nbB, nbB, muB, nuB, kuB, MULADD, pfA, latB,                &FFetch, &ifetch, &nfetch);   }/* * Save NB we've found */   sprintf(fnam, "res/%cNB", pre);   fp = fopen(fnam, "w");   fprintf(fp, "%d\n%d\n", 1, nbB);   fclose(fp);/* * Save best case parameters we have found */   CreateFinalSumm(pre, MULADD, pfA, latB, nbB, muB, nuB, kuB, FFetch, ifetch,                   nfetch, mfB);}void FindNC_0(char ta, char tb, char pre, int N, int mb, int nb, int kb,              int mu, int nu, int ku, int muladd, int pfA, int lat,              int FFetch, int ifetch, int nfetch){   int kuB=ku, latB=lat, lat0=lat, kb0=kb;   int i, j, k, csA=1, csB=1, csC=1, kmax;   double mf0, mf;   char fnam[128];   FILE *fp;   sprintf(fnam, "res/%cbest%c%c_%dx%dx%d", pre, ta, tb, mb, nb, kb);   if (FileExists(fnam)) /* default already exists */   {      GetInstLogFile(fnam, pre, &muladd, &pfA, &lat, &nb, &mu, &nu, &ku,                     &FFetch, &ifetch, &nfetch, &mf);      if (mf < 0.0) /* need to retime */      {         mf = mmcase(NULL, pre, "JIK", ta, tb, nb, nb, nb,                     nb, nb, nb, 0, 0, 0, mu, nu, ku, muladd, pfA, lat, 1,                     1, 1, csC, FFetch, ifetch, nfetch);         PutInstLogFile1(fnam, pre, muladd, pfA, lat, nb, mu, nu, ku,                         FFetch, ifetch, nfetch, mf);      }      return;   }   if (pre == 'c' || pre == 'z') csA = csB = csC = 2;   assert(N > 0);   if (kb == 0)   {      kb0 = 100000;      if ((mb*nb)/lat != lat) lat0 = GetGoodLat(muladd, kb0, mu, nu, 1, lat);   }   k = 1024 / (mu*nu);   for (kmax=4; kmax*kmax < k; kmax += 4);   if (pre == 'd' || pre == 's') kmax *= 2;   if (kmax >= N) kmax = N;   else if (kmax > N/2) kmax = N/2;   if (kb == 0) kuB = k = Mmin(ku,kmax);   else k = ku;/* * Find best non-cleanup case */   mf0 = mmcase(NULL, pre, "JIK", ta, tb, N, N, N, mb, nb, kb, 0, 0, 0,                mu, nu, k, muladd, pfA, lat0, 1, csA, csB, csC,                FFetch, ifetch, nfetch);   latB = lat0;/* * If kb is not known, try all available K unrollings; for large mu*nu*N * combinations, don't try maximal unrollings in order to avoid having * the compiler run out of space trying to optimize */   if (kb == 0)   {      for (k=1; k < kmax; k += 4)      {         if (k == 5) k = 4;         if (k > N/2) k = kmax;         j = k;         if (kb == 0) j = 1;         i = GetGoodLat(muladd, kb0, mu, nu, j, lat);         mf = mmcase(NULL, pre, "JIK", ta, tb, N, N, N, mb, nb, kb, 0, 0, 0,                     mu, nu, k, muladd, pfA, i, 1, csA, csB, csC,                     FFetch, ifetch, nfetch);         if (mf > mf0)         {            mf0 = mf;            kuB = k;            latB = i;         }      }   }/* * If K is known, try only the most common unrollings */   else   {      i = GetGoodLat(muladd, kb0, mu, nu, 1, lat);      mf = mmcase(NULL, pre, "JIK", ta, tb, N, N, N, mb, nb, kb, 0, 0, 0,                  mu, nu, 1, muladd, pfA, i, 1, csA, csB, csC,                  FFetch, ifetch, nfetch);      if (mf > mf0)      {         mf0 = mf;         kuB = 1;         latB = i;      }      i = GetGoodLat(muladd, kb0, mu, nu, 4, lat);      mf = mmcase(NULL, pre, "JIK", ta, tb, N, N, N, mb, nb, kb, 0, 0, 0,                  mu, nu, 4, muladd, pfA, i, 1, csA, csB, csC,                  FFetch, ifetch, nfetch);      if (mf > mf0)      {         mf0 = mf;         kuB = 4;         latB = i;      }      mf = mmcase(NULL, pre, "JIK", ta, tb, N, N, N, mb, nb, kb, 0, 0, 0,                  mu, nu, kb, muladd, pfA, lat, 1, csA, csB, csC,                  FFetch, ifetch, nfetch);      if (mf > mf0)      {         mf0 = mf;         kuB = kb;         latB = lat;      }   }/* * Try various latencies */   if (kb) i = kuB;   else i = 1;   for (k=2; k < 9; k++)   {      if (((mu*nu*i)/k)*k == mu*nu*i)      {         mf = mmcase(NULL, pre, "JIK", ta, tb, N, N, N, mb, nb, kb, 0, 0, 0,                     mu, nu, kuB, muladd, pfA, k, 1, csA, csB, csC,                     FFetch, ifetch, nfetch);         if (mf > mf0)         {            mf0 = mf;            latB = k;         }      }   }   fprintf(stdout, "BEST for %c%c_%dx%dx%d: mflop=%.2f\n",           ta, tb, mb, nb, kb, mf0);   fprintf(stdout,           "pre=%c ta=%c tb=%c nb=%d mu=%d nu=%d ku=%d muladd=%d lat=%d\n",           pre, ta, tb, nb, mu, nu, kuB, muladd, latB);   sprintf(fnam, "res/%cbest%c%c_%dx%dx%d", pre, ta, tb, mb, nb, kb);   fp = fopen(fnam, "w");   assert(fp);   PutInstLogFile(fp, muladd, pfA, latB, N, mu, nu, kuB,                  FFetch, ifetch, nfetch, mf0);   fclose(fp);}void FindNC0(char ta, char tb, char pre, int nb, int mu, int nu, int ku,             int muladd, int pfA, int lat, int FFetch, int ifetch, int nfetch){   FindNC_0(ta, tb, pre, nb, nb, nb, nb, mu, nu, ku, muladd, pfA, lat, FFetch,            ifetch, nfetch);   FindNC_0(ta, tb, pre, nb, 0, 0, nb, mu, nu, ku, muladd, pfA, lat, FFetch,            ifetch, nfetch);   FindNC_0(ta, tb, pre, nb, 0, 0, 0, mu, nu, ku, muladd, pfA, lat, FFetch,            ifetch, nfetch);}double NCcase(char pre, int nb, int mu, int nu, int ku, int ma, int pfA,              int lat, int ffetch, int ifetch, int nfetch){   double mf;   int ld=Mmax(1000,nb), cs=1;   char fnam[128];   if (pre == 'c' || pre == 'z') cs = 2;   do   {      sprintf(fnam, "res/%cNCNB%d_%d", pre, nb, ld);      mf = mmcase(fnam, pre, "JIK", 'N', 'N', nb, nb, nb, nb, nb, nb,                  -ld, nb, nb, mu, nu, ku, ma, pfA, lat, 1, cs, cs, cs,                  ffetch, ifetch, nfetch);      ld -= 10;   }   while (mf <= 0.0 && ld >= nb);   assert(mf > 0.0);   return(mf);}int FindNoCopyNB(char pre, int nb, int mu, int nu, int ku0, int muladd,                 int *prefA, int lat, int FFetch, int ifetch, int nfetch)/* * See if a smaller blocking factor is needed for no-copy */{   char fnam[128];   int i, ku, nbB=nb, csA=2, csB=2, csC=2, kuIsNB=0, pfA=(*prefA);   double mf, mfB, mf0;   const double dmul = 1.02;   FILE *fp;   if (ku0 == nb) kuIsNB = 1;   sprintf(fnam, "res/%cNCNB", pre);   if (!FileExists(fnam))   {/* *    Check both with and w/o prefetch, since no-copy prefetch different */      mfB = NCcase(pre, nb, mu, nu, ku0, muladd, pfA, lat,                   FFetch, ifetch, nfetch);      mf0 = NCcase(pre, nb, mu, nu, ku0, muladd, !pfA, lat,                   FFetch, ifetch, nfetch);      if (mf0 > mfB)      {         mfB = mf0;         pfA = !pfA;      }      mfB *= dmul;      mf0 = mfB;      for (i=nb-4; i >= 16; i -= 4)      {         if (kuIsNB) ku = i;         else ku = Mmin(i, ku0);         mf = NCcase(pre, i, mu, nu, ku, muladd, pfA, lat,                     FFetch, ifetch, nfetch);         if (1.2*mf < mfB) break; /* stop search after 20% slowdown */         if (nb%i == 0) mf *= dmul; /* give modest bonus to mults of nb */         if (mf > mfB)         {            mfB = mf;            nbB = i;         }      }/* *    For safety, check opposite of prefetch result wt new NB */      mf = NCcase(pre, nbB, mu, nu, ku, muladd, !pfA, lat,                  FFetch, ifetch, nfetch);      if (mf > mfB)      {         mfB = mf;         pfA = !pfA;      }      fp = fopen(fnam, "w");      assert(fp);      fprintf(fp, "%d\n", nbB);   }   else   /* If we know the correct NB, just try prefetch or not */   {      fp = fopen(fnam, "r");      *prefA = -1;      fscanf(fp, "%d\n", &nbB);      mf0 = mfB = -1.0;      ku = kuIsNB ? nbB : (Mmin(ku0,nbB));      mfB = NCcase(pre, nbB, mu, nu, ku, muladd, pfA, lat,                   FFetch, ifetch, nfetch);      mf  = NCcase(pre, nbB, mu, nu, ku, muladd, !pfA, lat,                   FFetch, ifetch, nfetch);      if (mf > mfB)

⌨️ 快捷键说明

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