mmsearch.c
来自「基于Blas CLapck的.用过的人知道是干啥的」· C语言 代码 · 共 2,188 行 · 第 1/5 页
C
2,188 行
fprintf(stdout,"\npre=%c, muladd=%d, lat=%d, pf=%d, nb=%d, mu=%d, nu=%d, ku=%d, mflop=%.2f\n", pre, MULADD, lat, pfA, NB, mu, nu, ku, t0); return(t0);}double mmcase0(char *nam, char pre, char *loopO, char ta, char tb, int M, int N, int K, int mb, int nb, int kb, int lda, int ldb, int ldc, int mu, int nu, int ku, int muladd, int pfA, int lat, int beta, int csA, int csB, int csC, int FFetch, int ifetch, int nfetch, char *mmnam){ char fnam[128], ln[512], bnam[16], casnam[128], mmcase[128]; int i, N0, lda2=lda, ldb2=ldb, ldc2=ldc; double mflop[NTIM], t0; FILE *fp; if (lda < 0) { lda2 = -lda; lda = 0; } if (ldb < 0) { ldb2 = -ldb; ldb = 0; } if (ldc < 0) { ldc2 = -ldc; ldc = 0; } if (mmnam) sprintf(mmcase, "mmucase mmrout=%s", mmnam); else sprintf(mmcase, "mmcase"); if (ifetch == -1 || nfetch == -1) { ifetch = mu+nu; nfetch = 1; } if (beta == 1) sprintf(bnam, "_b1"); else if (beta == -1) sprintf(bnam, "_bn1"); else if (beta == 0) sprintf(bnam, "_b0"); else sprintf(bnam, "_bX"); N0 = Mmax(M,N); if (N0 < K) N0 = K; if (ku > K) ku = K; else if (ku == -1) ku = K; if (nam) { strcpy(fnam, nam); sprintf(casnam, "casnam=%s", nam); } else { sprintf(fnam, "res/%c%smm%c%c%d_%dx%dx%d_%dx%dx%d_%dx%dx%d%s%s_%dx%d_%d_pf%d", pre, loopO, ta, tb, N0, mb, nb, kb, lda, ldb, ldc, mu, nu, ku, "_a1", bnam, muladd, lat, 1, pfA); casnam[0] = '\0'; } if (!FileExists(fnam)) { if (pre == 'c' || pre == 'z') sprintf(ln," make %s pre=%c loopO=%s ta=%c tb=%c M=%d N=%d K=%d mb=%d nb=%d kb=%d lda=%d ldb=%d ldc=%d lda2=%d ldb2=%d ldc2=%d mu=%d nu=%d ku=%d alpha=%d beta=%d muladd=%d lat=%d cleanup=%d csA=%d csB=%d csC=%d ff=%d if=%d nf=%d pfA=%d %s\n", mmcase,pre, loopO, ta, tb, M, N, K, mb, nb, kb, lda, ldb, ldc, lda2, ldb2, ldc2, mu, nu, ku, 1, beta, muladd, lat, 1, csA, csB, csC, FFetch, ifetch, nfetch, pfA, casnam); else sprintf(ln," make %s pre=%c loopO=%s ta=%c tb=%c M=%d N=%d K=%d mb=%d nb=%d kb=%d lda=%d ldb=%d ldc=%d lda2=%d ldb2=%d ldc2=%d mu=%d nu=%d ku=%d alpha=%d beta=%d muladd=%d lat=%d cleanup=%d ff=%d if=%d nf=%d pfA=%d %s\n", mmcase, pre, loopO, ta, tb, M, N, K, mb, nb, kb, lda, ldb, ldc, lda2, ldb2, ldc2, mu, nu, ku, 1, beta, muladd, lat, 1, FFetch, ifetch, nfetch, pfA, casnam); fprintf(stderr, "%s:\n",ln); if (system(ln) != 0) {/* * User cases, and large leading dimensions can fail to run */ if (mmnam) return(-1.0); /* user cases can fail to compile */ if (lda2 != lda || ldb2 != ldb || ldc2 != ldc) return(-1); fprintf(stderr, "Error in command: %s", ln); sprintf(ln, "rm -f %s\n", fnam); system(ln); exit(-1); } } fp = fopen(fnam, "r"); if (!fp) fprintf(stderr, "ERROR: can't find file=%s\n", fnam); assert(fp); for (i=0; i != NTIM; i++) { assert(fscanf(fp, "%lf", &mflop[i]) == 1); } fclose(fp); t0 = GetAvg(NTIM, TOLERANCE, mflop); if (t0 == -1.0) { fprintf(stderr, "case=%s: rerun with higher reps; variation exceeds tolerence\n", fnam); sprintf(ln, "rm -f %s\n", fnam); system(ln); exit(-1); } fprintf(stdout,"\n pre=%c, loopO=%s, ta=%c tb=%c, mb=%d, nb=%d, kb=%d, lda=%d, ldb=%d, ldc=%d\n", pre, loopO, ta, tb, mb, nb, kb, lda, ldb, ldc); fprintf(stdout, " mu=%d, nu=%d, ku=%d, muladd=%d, lat=%d ====> mflop=%f\n", mu, nu, ku, muladd, lat, t0); return(t0);}double mmucase(int ifile, char pre, int nb, int muladd, int lat, int mu, int nu, int ku, char *fnam){ char fout[64]; int iff; sprintf(fout, "res/%cuser%d", pre, ifile); if (mu == 1 && nu == 1) iff = 1; else iff = mu + nu; return(mmcase0(fout, pre, "JIK", 'T', 'N', nb, nb, nb, nb, nb, nb, nb, nb, 0, mu, nu, ku, muladd, 0, lat, 1, 1, 1, 2, 0, iff, 1, fnam));}enum CW {CleanM=0, CleanN=1, CleanK=2, CleanNot=3};double mmclean(char pre, enum CW which, char *loopO, char ta, char tb, int M, int N, int K, int mb, int nb, int kb, int lda, int ldb, int ldc, int mu, int nu, int ku, int muladd, int pfA, int lat, int beta, int csA, int csB, int csC, int FFetch, int ifetch, int nfetch){ char nam[128]; char cwh[3] = {'M', 'N', 'K'}; sprintf(nam, "res/%cClean%c_%dx%dx%d", pre, cwh[which], M, N, K); return(mmcase0(nam, pre, loopO, ta, tb, M, N, K, mb, nb, kb, lda, ldb, ldc, mu, nu, ku, muladd, pfA, lat, beta, csA, csB, csC, FFetch, ifetch, nfetch, NULL));}double mmcase(char *nam, char pre, char *loopO, char ta, char tb, int M, int N, int K, int mb, int nb, int kb, int lda, int ldb, int ldc, int mu, int nu, int ku, int muladd, int pfA, int lat, int beta, int csA, int csB, int csC, int FFetch, int ifetch, int nfetch){ return(mmcase0(nam, pre, loopO, ta, tb, M, N, K, mb, nb, kb, lda, ldb, ldc, mu, nu, ku, muladd, pfA, lat, beta, csA, csB, csC, FFetch, ifetch, nfetch, NULL));}int GetGoodLat(int MULADD, int kb, int mu, int nu, int ku, int lat){ int slat, blat, i, ii = mu*nu*ku; if (MULADD) return(lat); if ( (lat > 1) && (kb > ku) && ((ii/lat)*lat != ii) ) /* lat won't work */ { for (i=lat; i; i--) if ( (ii/i) * i == ii ) break; slat = i; for (i=lat; i < MAXLAT; i++) if ( (ii/i) * i == ii ) break; blat = i; if ( (ii/blat)*blat != ii ) blat = slat; if (slat < 2) lat = blat; else if (lat-slat < blat-lat) lat = slat; else lat = blat; } return(lat);}void FindMUNU(int muladd, /* 0: use separate multiply and add inst */ int lat, /* pipe len for muladd=0 */ int nr, /* # of registers available */ int FullTest, /* 0: use shortcut if available */ int *MU, /* suggested MU */ int *NU) /* suggested NU *//* * Find near-square muxnu using nr registers or less */{ int j, mu, nu; if (nr < 1) { *MU = lat; *NU = 1; return; } if (muladd) j = nr; else j = nr - lat; if (j < 3) mu = nu = 1; else {/* * For x86, two-operand assembler means 1-D case almost certainly best */ #ifdef TWO_OP_ASM if (!FullTest) { if (lat > 2) mu = lat; else mu = 4; nu = 1; } else { #endif mu = j + 1; for (nu=1; nu*nu < mu; nu++); if (nu*nu > mu) nu -= 2; else nu--; if (nu < 1) mu = nu = 1; else { mu = (nr-nu) / (1+nu); if (mu < 1) mu = 1; } if (mu < nu) { j = mu; mu = nu; nu = j; } #ifdef TWO_OP_ASM } #endif } *MU = mu; *NU = nu;}void PutInstLogLine(FILE *fp, int muladd, int pfA, int lat, int nb, int mu, int nu, int ku, int ForceFetch, int ifetch, int nfetch, double mflop){ fprintf(fp, "%6d %3d %4d %3d %3d %3d %3d %5d %5d %5d %7.2lf\n", muladd, lat, pfA, nb, mu, nu, ku, ForceFetch, ifetch, nfetch, mflop);}void PutInstLogFile(FILE *fp, int muladd, int pfA, int lat, int nb, int mu, int nu, int ku, int ForceFetch, int ifetch, int nfetch, double mflop){ fprintf(fp, "MULADD LAT PREF NB MU NU KU FFTCH IFTCH NFTCH MFLOP\n"); PutInstLogLine(fp, muladd, pfA, lat, nb, mu, nu, ku, ForceFetch, ifetch, nfetch, mflop);}void PutInstLogFile1(char *fnam, char pre, int muladd, int pfA, int lat, int nb, int mu, int nu, int ku, int ForceFetch, int ifetch, int nfetch, double mflop){ FILE *fp; fp = fopen(fnam, "w"); assert(fp); PutInstLogFile(fp, muladd, pfA, lat, nb, mu, nu, ku, ForceFetch, ifetch, nfetch, mflop); fclose(fp);}void GetInstLogLine(FILE *fp, int *muladd, int *pfA, int *lat, int *nb, int *mu, int *nu, int *ku, int *ForceFetch, int *ifetch, int *nfetch, double *mflop){ assert(fscanf(fp, " %d %d %d %d %d %d %d %d %d %d %lf\n", muladd, lat, pfA, nb, mu, nu, ku, ForceFetch, ifetch, nfetch, mflop) == 11);}void GetInstLogFile(char *nam, char pre, int *muladd, int *pfA, int *lat, int *nb, int *mu, int *nu, int *ku, int *ForceFetch, int *ifetch, int *nfetch, double *mflop){ char ln[128]; FILE *fp; fp = fopen(nam, "r"); if (fp == NULL) fprintf(stderr, "file %s not found!!\n\n", nam); assert(fp); fgets(ln, 128, fp); GetInstLogLine(fp, muladd, pfA, lat, nb, mu, nu, ku, ForceFetch, ifetch, nfetch, mflop); fclose(fp);}void CreateFinalSumm(char pre, int muladd, int pfA, int lat, int nb, int mu, int nu, int ku, int Ff, int If, int Nf, double gmf){ char ln[64], auth[65]; FILE *fp, *fp0; int icase, unb; double umf; sprintf(ln, "res/%cMMRES", pre); fp = fopen(ln, "w"); PutInstLogFile(fp, muladd, pfA, lat, nb, mu, nu, ku, Ff, If, Nf, gmf); sprintf(ln, "res/%cuMMRES", pre); fp0 = fopen(ln, "r"); assert(fp0); assert(fgets(ln, 64, fp0)); assert(fscanf(fp0, " %d %d %lf \"%[^\"]\" \"%[^\"]", &icase, &unb, &umf, ln, auth) == 5); fclose(fp0); fprintf(fp, "\nICASE NB MFLOP ROUT AUTHOR\n"); fprintf(fp, "%5d %3d %8.2f \"%.63s\" \"%.63s\"\n", icase, unb, umf, ln, auth); fclose(fp);}void FindFetch(char ta, char tb, char pre, int mb, int nb, int kb, int mu, int nu, int ku, int muladd, int pfA, int lat, int *FFetch0, int *ifetch0, int *nfetch0)/* * See what fetch patterns are appropriate */{ char fnam[128]; const int nelts = mu+nu; int csA=1, csB=1, csC=1, nleft, i, j; int ifetch = mu+nu, nfetch = 1; double mf, mf0; if (pre == 'c' || pre == 'z') csC = 2; mf0 = mmcase(NULL, pre, "JIK", ta, tb, mb, nb, kb, mb, nb, kb, kb, kb, 0, mu, nu, ku, muladd, pfA, lat, 0, csA, csB, csC, 0, ifetch, nfetch); for (i=2; i < nelts; i++) { nleft = nelts - i; for (j=1; j <= nleft; j++) { sprintf(fnam, "res/%cMMfetch%d_%d", pre, i, j); mf = mmcase(fnam, pre, "JIK", ta, tb, mb, nb, kb, mb, nb, kb, kb, kb, 0, mu, nu, ku, muladd, pfA, lat, 0, csA, csB, csC, 0, i, j); if (mf > mf0) { mf = mf0; ifetch = i; nfetch = j; } } }/* * See if prefetching good idea for beta=0 case */ sprintf(fnam, "res/%cMM_b0", pre); mf0 = mmcase(fnam, pre, "JIK", ta, tb, mb, nb, kb, mb, nb, kb, kb, kb, 0, mu, nu, ku, muladd, pfA, lat, 0, csA, csB, csC, 0, ifetch, nfetch); sprintf(fnam, "res/%cMM_b0_pref", pre); mf = mmcase(fnam, pre, "JIK", ta, tb, mb, nb, kb, mb, nb, kb, kb, kb, 0, mu, nu, ku, muladd, pfA, lat, 0, csA, csB, csC, 1, ifetch, nfetch); *FFetch0 = (mf > mf0); *ifetch0 = ifetch; *nfetch0 = nfetch; fprintf(stdout, "\n\nFORCEFETCH=%d, IFETCH = %d, NFETCH = %d\n\n", *FFetch0, *ifetch0, *nfetch0);}void kucases(char pre, int muladd, int nb, int mu, int nu, int pfA, int LAT, int Fku, double *mfB, int *nbB, int *muB, int *nuB, int *kuB, int *latB){ double mf; int ku, lat; if (Fku == -1 || !Fku) ku = nb; else ku = Fku; if (ku != nb) lat = GetGoodLat(muladd, nb, mu, nu, ku, LAT); else lat = LAT; mf = mms_case(pre, muladd, nb, mu, nu, ku, pfA, lat); if (mf > *mfB) { *mfB = mf; *nbB = nb; *muB = mu; *nuB = nu; *kuB = ku; *latB = lat; } if (!Fku) { lat = GetGoodLat(muladd, nb, mu, nu, 1, LAT); mf = mms_case(pre, muladd, nb, mu, nu, 1, pfA, lat); if (mf > *mfB) { *mfB = mf; *nbB = nb; *muB = mu; *nuB = nu; *kuB = 1; *latB = lat; } }}void searchnu_nu(char pre, int nb, int maxreg, int Fku, int muladd, int pfA, int LAT, int NO1D, double *mfB, int *nbB, int *muB, int *nuB, int *kuB, int *latB)/* * For large number of registers, search only the near-square cases */{ int i, k, mu, nu, lat; double mf; for (k=16; k <= maxreg; k += 4) { for (nu=1; nu*nu < k; nu++); nu--; mu = k / nu; kucases(pre, muladd, nb, mu, nu, pfA, LAT, Fku, mfB, nbB, muB, nuB, kuB, latB); if (mu != nu) /* reverse them */ kucases(pre, muladd, nb, nu, mu, pfA, LAT, Fku, mfB, nbB, muB, nuB, kuB, latB); }}#ifdef TWO_OP_ASMvoid searchmu_nu(char pre, int nb, int maxreg, int Fku, int muladd, int pfA, int LAT, int NO1D, double *mfB, int *nbB, int *muB, int *nuB, int *kuB, int *latB)/* * For two-op assembly machines, search only 1-D register blockings */{ int i, n, nreg=Mmin(nb, maxreg); n = (nreg-1)/2; for (i=1; i <= n; i++) { kucases(pre, muladd, nb, i, 1, pfA, LAT, Fku, mfB, nbB, muB, nuB, kuB, latB); if (i != 1) kucases(pre, muladd, nb, 1, i, pfA, LAT, Fku, mfB, nbB, muB, nuB, kuB, latB); }}/* * > 2 op machine can benefit from 2-D register blockings */#elsevoid searchmu_nu(char pre, int nb, int maxreg, int Fku, int muladd, int pfA, int LAT, int NO1D, double *mfB, int *nbB, int *muB, int *nuB, int *kuB, int *latB)
⌨️ 快捷键说明
复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?