emit_mm.c
来自「基于Blas CLapck的.用过的人知道是干啥的」· C语言 代码 · 共 1,670 行 · 第 1/5 页
C
1,670 行
/* * fetchC : fetches mu*nu elts of C, & applies beta * If LoopOrder is IJK, we will pass B in as A, so we need to transpose C */void fetchC(FILE *fpout, char *spc, enum ATLAS_LOOP_ORDER LoopOrder, int ForceFetch, /* fetch C even if beta==0? */ int mu, /* unroll of inner loop */ int nu, /* unroll of outer loop */ int offC, /* offset to start of C */ char *pC, /* varnam of pointer to C */ char *rC, /* register name for elts of C */ int mul, /* stride in elts of C in column (1=real, 2=cplx) */ int ldc, /* if 0: use ptrs for cols of C; else row stride */ int alpha, int beta, int fetch, /* 0: just do lame pref, don't actually fetch C */ char *reg) /* name of unrelated register for beta */{ int i, j;/* * If we aren't fetching C, only thing to do is ForceFetch, and we might * as well act as if beta = 0 */ if (!fetch) { if (!ForceFetch) return; beta = 0; } if (ForceFetch && beta == 0) /* lame-ass prefetch on C */ { fprintf(fpout, "/*\n *%s Feeble prefetch of C\n */\n", spc); for (j=0; j < nu; j++) { fprintf(fpout, "%s %s%d_%d = ", spc, rC, 0, j); if (ldc) { if (offC+j) fprintf(fpout, "%s0[%d];\n", pC, offC+(j*ldc)*mul); else fprintf(fpout, "*%s0;\n", pC); } else { if (offC) fprintf(fpout, "%s%d[%d];\n", pC, j, offC); else fprintf(fpout, "*%s%d;\n", pC, j); } } if (!fetch) return; } else if ( (beta && alpha != SAFE_ALPHA) ) { if (alpha != 1) fprintf(fpout, "%s %s = BetaAlpha;\n", spc, reg); else if (beta != 1 && beta != -1 && beta != 0) fprintf(fpout, "%s %s = beta;\n", spc, reg); for (j=0; j < nu; j++) { for (i=0; i < mu; i++) { if (LoopOrder == AtlasJIK) fprintf(fpout, "%s %s%d_%d = ", spc, rC, i, j); else if (LoopOrder == AtlasIJK) fprintf(fpout, "%s %s%d_%d = ", spc, rC, j, i); if (ldc) { if (offC+i+j) fprintf(fpout, "%s0[%d];\n", pC, offC+(j*ldc+i)*mul); else fprintf(fpout, "*%s0;\n", pC); } else { if (offC+i) fprintf(fpout, "%s%d[%d];\n", pC, j, offC+i*mul); else fprintf(fpout, "*%s%d;\n", pC, j); } if (beta == -1) fprintf(fpout, "%s %s%d_%d = -%s%d_%d;\n", spc, rC, i, j, rC, i, j); else if (beta != 1 && beta != 0 || alpha != 1) fprintf(fpout, "%s %s%d_%d *= %s;\n", spc, rC, i, j, reg); } } } if (beta == 0 || alpha == SAFE_ALPHA) { fprintf(fpout, "%s ", spc); for (j=0; j < nu; j++) for (i=0; i < mu; i++) fprintf(fpout, "%s%d_%d = ", rC, i, j); fprintf(fpout, "0.0;\n"); }}void IncPtrs(FILE *fpout, char *spc, int np, /* number of pointers to increment */ char var, /* variable: A, B, C */ char loop) /* which loop: n, m, k */{ int i; for (i=0; i < np; i++) fprintf(fpout, "%s p%c%d += inc%c%c;\n", spc, var, i, var, loop);}void Cass(FILE *fpout, char *spc, enum ATLAS_LOOP_ORDER LoopOrder, int LoadC, /* 1 load C & apply beta before assignment */ int mu, /* unroll of inner loop */ int nu, /* unroll of outer loop */ int alpha, int beta, int offC, /* offset to start of C */ char *rA, /* register name for elts of A */ char *rB, /* register name for elts of B */ char *pC, /* varnam of pointer to C */ char *rC, /* register name for elts of C */ int mulC, /* stride in elts of C in column (1=real, 2=cplx) */ int ldc, /* if 0: use ptrs for cols of C; else row stride */ char *incC) /* increment for this loop; if NULL don't use */{ int i, j; char *cp, calpha[32], Cderef[32]; int AlphaReg = (nu > 1); if (!beta) LoadC = 0; /* never load C if beta == 0.0 */ if (LoadC && alpha == SAFE_ALPHA) alpha = SAFE_ALPHA - 1; if (alpha == SAFE_ALPHA) { if (LoopOrder == AtlasIJK) { cp = rA; rA = rB; rB = cp; if (nu > 1) sprintf(calpha, "%s1", rB); else sprintf(calpha, "alpha"); } else { if (nu > 1) sprintf(calpha, "%s1", rB); else sprintf(calpha, "alpha"); } fprintf(fpout, "%s %s0 = beta;\n", spc, rB); if (AlphaReg) fprintf(fpout, "%s %s = alpha;\n", spc, calpha); for (j=0; j < nu; j++) { for (i=0; i < mu; i++) { fprintf(fpout, "%s %s%d_%d *= %s;\n", spc, rC, i, j, calpha); if (ldc) fprintf(fpout, "%s %s%d = %s0[%d];\n", spc, rA, i, pC, mulC*(j*ldc+i)); else fprintf(fpout, "%s %s%d = %s%d[%d];\n", spc, rA, i, pC, j, i*mulC); fprintf(fpout, "%s %s%d_%d += %s0 * %s%d;\n", spc, rC, i, j, rB, rA, i); } } } else if (alpha != 1) { fprintf(fpout, "%s %s0 = alpha;\n", spc, rB); for (i=0; i < mu; i++) { for (j=0; j < nu; j++) { if (LoopOrder == AtlasJIK) fprintf(fpout, "%s %s%d_%d *= %s0;\n", spc, rC, i, j, rB); else if (LoopOrder == AtlasIJK) fprintf(fpout, "%s %s%d_%d *= %s0;\n", spc, rC, j, i, rB); } } } if (LoadC && beta != 0 && beta != 1) fprintf(fpout, "%s %s0 = beta;\n", spc, rB); for (j=0; j < nu; j++) { for (i=0; i < mu; i++) { if (ldc) { if (offC+j+i) sprintf(Cderef, "%s0[%d]", pC, offC+(ldc*j+i)*mulC); else sprintf(Cderef, "*%s0", pC); } else { if (i) sprintf(Cderef, "%s%d[%d]", pC, j, i*mulC); else sprintf(Cderef, "*%s%d", pC, j); } fprintf(fpout, "%s %s", spc, Cderef); if (LoadC) { if (beta == 1) fprintf(fpout, " += "); else if (beta == 0) fprintf(fpout, " = "); else fprintf(fpout, " = %s * %s0 + ", Cderef, rB); if (LoopOrder == AtlasJIK) fprintf(fpout, "%s%d_%d;\n", rC, i, j); else if (LoopOrder == AtlasIJK) fprintf(fpout, "%s%d_%d;\n", rC, j, i); } else { if (LoopOrder == AtlasJIK) fprintf(fpout, " = %s%d_%d;\n", rC, i, j); else if (LoopOrder == AtlasIJK) fprintf(fpout, " = %s%d_%d;\n", rC, j, i); } } }}#define RegCallSeq 1void CallMM(FILE *fpout, char *spc, char pre, char *loopstr, int CleanUp, enum ATLAS_TRANS TA, enum ATLAS_TRANS TB, int M, int N, int K, int mu, int nu, int ku, int alpha, int beta, int lda, int ldb, int ldc, int Mb, int Nb, int Kb, char *cA, char *cB, char *cC, char *cM, char *cN, char *cK){ char cTA='N', cTB='N'; if (TA == AtlasTrans) cTA = 'T'; else if (TA == AtlasConjTrans) cTA = 'C'; if (TB == AtlasTrans) cTB = 'T'; else if (TB == AtlasConjTrans) cTB = 'C'; if (CleanUp) fprintf(fpout, "%s ATL_%c%s%dx%dx%d%c%c%dx%dx%d", spc, pre, loopstr, M, N, K, cTA, cTB, mu, nu, ku); else fprintf(fpout, "%s ATL_%c%s%dx%dx%d%c%c%dx%dx%d", spc, pre, loopstr, M, N, K, cTA, cTB, lda, ldb, ldc); if (alpha == 1) fprintf(fpout, "_a1"); else if (alpha == -1) fprintf(fpout, "_an1"); else if (alpha == SAFE_ALPHA) fprintf(fpout, "_aXX"); else fprintf(fpout, "_aX"); if (beta == 1 || beta == 0) fprintf(fpout, "_b%d(", beta); else if (beta == -1) fprintf(fpout, "_bn1("); else fprintf(fpout, "_bX("); if (M) fprintf(fpout, "%d, ", M); else fprintf(fpout, "%s, ", cM); if (N) fprintf(fpout, "%d, ", N); else fprintf(fpout, "%s, ", cN); if (K) fprintf(fpout, "%d, ", K); else fprintf(fpout, "%s, ", cK); if ((alpha != 1 && alpha != -1) || RegCallSeq) fprintf(fpout, "alpha, "); fprintf(fpout, "%s, ", cA); if (!lda || RegCallSeq) fprintf(fpout, "lda, "); fprintf(fpout, "%s, ", cB); if (!ldb || RegCallSeq) fprintf(fpout, "ldb, "); if ( (beta != 1 && beta != 0 && beta != -1) || RegCallSeq) fprintf(fpout, "beta, "); fprintf(fpout, "%s", cC); if (!ldc || RegCallSeq) fprintf(fpout, ", ldc"); fprintf(fpout, ");\n");}void MMDeclare(FILE *fpout, char *spc, char pre, char *type, char *decmod, char *loopstr, enum ATLAS_TRANS TA, enum ATLAS_TRANS TB, int M, int N, int K, int mu, int nu, int ku, int alpha, int beta, int lda, int ldb, int ldc, int pfA){ char cTA='N', cTB='N'; if (TA == AtlasTrans) cTA = 'T'; else if (TA == AtlasConjTrans) cTA = 'C'; if (TB == AtlasTrans) cTB = 'T'; else if (TB == AtlasConjTrans) cTB = 'C';/* * For cleanup, put unrolling in name to distinguish from original function. * For regular function, encode leading dimensions instead */ if (decmod[0] == '\0') fprintf(fpout, "%svoid ATL_%c%s%dx%dx%d%c%c%dx%dx%d", spc, pre, loopstr, M, N, K, cTA, cTB, lda, ldb, ldc); else fprintf(fpout, "%s%svoid ATL_%c%s%dx%dx%d%c%c%dx%dx%d", spc, decmod, pre, loopstr, M, N, K, cTA, cTB, mu, nu, ku); if (alpha == 1) fprintf(fpout, "_a1"); else if (alpha == -1) fprintf(fpout, "_an1"); else if (alpha == SAFE_ALPHA) fprintf(fpout, "_aXX"); else fprintf(fpout, "_aX"); if (beta == 1 || beta == 0) fprintf(fpout, "_b%d", beta); else if (beta == -1) fprintf(fpout, "_bn1"); else fprintf(fpout, "_bX"); fprintf(fpout, "\n ("); if (!M || RegCallSeq) fprintf(fpout, "const int M, "); if (!N || RegCallSeq) fprintf(fpout, "const int N, "); if (!K || RegCallSeq) fprintf(fpout, "const int K, "); if ( (alpha != 1 && alpha != -1) || RegCallSeq) fprintf(fpout, "const %s alpha, ", type); fprintf(fpout, "const %s * ATL_RESTRICT A, ", type); if (!lda || RegCallSeq) fprintf(fpout, "const int lda, "); fprintf(fpout, "const %s * ATL_RESTRICT B, ", type); if (!ldb || RegCallSeq) fprintf(fpout, "const int ldb, "); if (beta != 1 && beta != -1 && beta != 0 || RegCallSeq) fprintf(fpout, "const %s beta, ", type); fprintf(fpout, "%s * ATL_RESTRICT C", type); if (!ldc || RegCallSeq) fprintf(fpout, ", const int ldc"); fprintf(fpout, ")\n"); fprintf(fpout, "/*\n * matmul with TA=%c, TB=%c, MB=%d, NB=%d, KB=%d, \n", cTA, cTB, M, N, K); fprintf(fpout, " * lda=%d, ldb=%d, ldc=%d, mu=%d, nu=%d, ku=%d, pf=%d\n", lda, ldb, ldc, mu, nu, ku, pfA); fprintf(fpout, " * Generated by ATLAS/tune/blas/gemm/emit_mm.c (3.8.0)\n"); fprintf(fpout, " */\n");}static int ncucases=0;static int cucases[64][7];/* * For every operand involved in a given loop, we pass: * ldx : if 0, use pointer for each column, and moving between columns is * accomplished by incrementing by constant passed in string incX * else: use only 1 pointer, increment by ldx move between columns */void emit_mm(FILE *fpout, char *spc, /* indentation string */ char pre, char *type, char *decmod, /* routine declaration modifier (eg. static) */ enum ATLAS_LOOP_ORDER LoopOrder, enum ATLAS_TRANS TA, enum ATLAS_TRANS TB, int CleanUp, /* 1 : issue cleanup code, 0: do not */ int muladd, /* 0: separate mult & add, ELSE: combined muladd */ int prefA, /* 0: don't prefetch nxt blk of A */ int lat, /* pipeline length */ int ForceFetch, int ifetch, /* number of initial fetches to perform */ int nfetch, /* # of fetches to perform for every flop */ int M, /* if 0, M is arbitrary, else M is len of M-loop */ int N, /* if 0, N is arbitrary, else N is len of N-loop */ int K,
⌨️ 快捷键说明
复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?