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 + -
显示快捷键?