emit_mm.c
来自「基于Blas CLapck的.用过的人知道是干啥的」· C语言 代码 · 共 1,670 行 · 第 1/5 页
C
1,670 行
for (i=1; i < mu; i++) fprintf(fpout, "%s %s%d += %s;\n", spc, pA, i, incAk); } if (!incpB) /* need to increment the B pointers ourselves */ { fprintf(fpout, "%s %s0 += %s;\n", spc, pB, incBk); if (rowB && !ldb) for (i=1; i < nu; i++) fprintf(fpout, "%s %s%d += %s;\n", spc, pB, i, incBk); } spc += 3; fprintf(fpout, "%s }\n", spc);/* * K-loop cleanup */ if (K) /* known cleanup */ { if (K%ku) regKunroll_ma(fpout, spc, LoopOrder, 0, ifetch, nfetch, rA, rB, rC, pA, pB, mu, nu, K%ku, &offA, &offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); sprintf(kadj, "%d", K%ku); } else if (ku > 1) /* cleanup between 0-(ku-1) */ { fprintf(fpout, "%s switch(k = (K-Kb))\n%s {\n", spc, spc); for (i=1; i < ku; i++) { fprintf(fpout, "%s case %d:\n", spc, i); regKunroll_ma(fpout, spc-3, LoopOrder, 0, ifetch, nfetch, rA, rB, rC, pA, pB, mu, nu, i, &offA, &offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); fprintf(fpout, "%s break;\n", spc); } fprintf(fpout, "%s case 0: ;\n%s }\n", spc, spc); sprintf(kadj, "k"); } if ( (K && K%ku) || ku > 1) { if (incpA) { fprintf(fpout, "%s %s0 -= %s*%s;\n", spc, pA, kadj, incAk); if (lda == 0 && TA != AtlasNoTrans) for (i=1; i < mu; i++) fprintf(fpout, "%s %s%d -= %s*%s;\n", spc, pA, i, kadj, incAk); } if (incpB) { fprintf(fpout, "%s %s0 -= %s*%s;\n", spc, pB, kadj, incBk); if (rowB && !ldb) for (i=1; i < nu; i++) fprintf(fpout, "%s %s%d -= %s*%s;\n", spc, pB, i, kadj, incBk); } }}/* * regKunroll: K-loop unrolling for separate multiply and add instructions. * Unrolls k loop by ku, with mu unroll along inner matrix & nu along outer. * If actual inner matrix is B instead of A, pass B's data in A and vice versa, * and assign pC[i][j] = rCji. * Assumes lat flops of this iteration done previously, lat of next done here, * Present flop to be done is (iop,jop) of the mu*nu to be done. * assumes ifetch has enough data for all lat operations !!!! */int regKunroll(FILE *fpout, char *spc, /* indentation string */ enum ATLAS_LOOP_ORDER LoopOrder, int ifetch, /* number of initial fetches to perform */ int nfetch, /* # of fetches to perform for every flop */ int lat, /* skew of K-loop */ int STARTUP,/* 0: not starting pipeline, else doing so */ int Asg1stC,/* 1: first C update gets =, instead of += */ char *rA, /* varnam for registers of inner matrix */ char *rB, /* varnam for registers of outer matrix */ char *rC, /* varnam for registers of C */ char *pA, /* varnam for pointer(s) to inner matrix */ char *pB, /* varnam for pointer(s) to outer matrix */ int mu, /* unrolling along inner loop */ int nu, /* unrolling along outer loop */ int ku, /* unrolling along k-loop (innermost) */ int *offA, /* offset to first elt of this block */ int *offB, /* offset to first elt of this block */ int lda, /* row stride; if 0, row stride is arbitrary */ int ldb, /* row stride; if 0, row stride is arbitrary */ int mulA, /* col stride; 1: real 2: cplx */ int mulB, /* col stride; 1: real 2: cplx */ int incpA, /* Increment A for every K iteration? */ int incpB, /* Increment B for every K iteration? */ char *incAk, /* if !rowA, k-loop increment for ptrs */ char *incBk, /* if !rowB, k-loop increment for ptrs */ int rowA, /* if 0, fetch within col, else fetch within row */ int rowB, /* if 0, fetch within col, else fetch within row */ int *ia, /* elt of inner matrix to be fetched */ int *ib, /* elt of outer matrix to be fetched */ int *iop0, /* 1st operand of inner matrix to use */ int *jop0) /* 1st operand of outer matrix to use */{ int i, j, k, h=0; int iop = *iop0, jop = *jop0; if (LoopOrder == AtlasIJK) return(regKunroll(fpout, spc, AtlasJIK, ifetch, nfetch, lat, STARTUP, Asg1stC, rB, rA, rC, pB, pA, nu, mu, ku, offB, offA, ldb, lda, mulB, mulA, incpB, incpA, incBk, incAk, rowB, rowA, ib, ia, jop0, iop0)); if (STARTUP) fprintf(fpout, "/*\n *%s Start pipeline\n */\n", spc);/* * If we have not fetched any data for this iteration yet, do so */ if (STARTUP && *ia == 0 && *ib == 0) opfetch(fpout, spc, ifetch, rA, rB, pA, pB, mu, nu, *offA, *offB, lda, ldb, mulA, mulB, rowA, rowB, ia, ib);/* * One iteration of the ku-unrolled K loop */ for (k=0; k < ku; k++) { for (j=0; j < nu; j++) { for (i=0; i < mu; i++) { if (!STARTUP) { if (Asg1stC && !k) fprintf(fpout, "%s %s%d_%d = m%d;\n", spc, rC, i, j, h); else fprintf(fpout, "%s %s%d_%d += m%d;\n", spc, rC, i, j, h); } fprintf(fpout, "%s m%d = %s%d * %s%d;\n", spc, h, rA, iop, rB, jop); if (++iop == mu) { iop = 0; if (++jop == nu) /* used all this iteration's data */ { jop = 0; incABk(fpout, spc, pA, pB, mu, nu, offA, offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); *ia = *ib = 0; opfetch(fpout, spc, ifetch, rA, rB, pA, pB, mu, nu, *offA, *offB, lda, ldb, mulA, mulB, rowA, rowB, ia, ib); } } opfetch(fpout, spc, nfetch, rA, rB, pA, pB, mu, nu, *offA, *offB, lda, ldb, mulA, mulB, rowA, rowB, ia, ib); if (++h == lat) { if (STARTUP) { fprintf(fpout, "\n"); *iop0 = iop; *jop0 = jop; return(lat); } h = 0; } } } } *iop0 = iop; *jop0 = jop; return(h);}/* * regKdrain: drains the pipe by explicitly unrolling last K iteration */void regKdrain(FILE *fpout, char *spc, /* indentation string */ enum ATLAS_LOOP_ORDER LoopOrder, int ifetch, /* number of initial fetches to perform */ int nfetch, /* # of fetches to perform for every flop */ int lat, /* skew of K-loop */ char *rA, /* varnam for registers of inner matrix */ char *rB, /* varnam for registers of outer matrix */ char *rC, /* varnam for registers of C */ char *pA, /* varnam for pointer(s) to inner matrix */ char *pB, /* varnam for pointer(s) to outer matrix */ int mu, /* unrolling along inner loop */ int nu, /* unrolling along outer loop */ int ku, /* unrolling along k-loop (innermost) */ int *offA, /* offset to first elt of this block */ int *offB, /* offset to first elt of this block */ int lda, /* row stride; if 0, row stride is arbitrary */ int ldb, /* row stride; if 0, row stride is arbitrary */ int mulA, /* col stride; 1: real 2: cplx */ int mulB, /* col stride; 1: real 2: cplx */ int incpA, /* Increment A for every K iteration? */ int incpB, /* Increment B for every K iteration? */ char *incAk, /* if !rowA, k-loop increment for ptrs */ char *incBk, /* if !rowB, k-loop increment for ptrs */ int rowA, /* if 0, fetch within col, else fetch within row */ int rowB, /* if 0, fetch within col, else fetch within row */ int *ia, /* elt of inner matrix to be fetched */ int *ib, /* elt of outer matrix to be fetched */ int iop, /* 1st operand of inner matrix to use */ int jop, /* 1st operand of outer matrix to use */ int h) /* */{ int i, j, k=0; int REGFETCH=1; if (LoopOrder == AtlasIJK) { regKdrain(fpout, spc, AtlasJIK, ifetch, nfetch, lat, rB, rA, rC, pB, pA, nu, mu, ku, offB, offA, ldb, lda, mulB, mulA, incpB, incpA, incBk, incAk, rowB, rowA, ib, ia, jop, iop, h); return; }/* * Drain part of pipe where we are still doing multiplies. Once iop and jop * reach mu/nu, we stop doing fetches */ fprintf(fpout, "/*\n *%s Drain pipe on last iteration of K-loop\n */\n", spc); for (j=0; j < nu; j++) { for (i=0; i < mu; i++) { fprintf(fpout, "%s %s%d_%d += m%d;\n", spc, rC, i, j, h); if (iop < mu || jop < nu) { fprintf(fpout, "%s m%d = %s%d * %s%d;\n", spc, h, rA, iop, rB, jop); if (++iop == mu) { iop = 0; if (++jop == nu) /* used all this iteration's data */ { REGFETCH = 0; iop = mu; incABk(fpout, spc, pA, pB, mu, nu, offA, offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); } } } else k++; if (++h == lat) h = 0; if (REGFETCH) opfetch(fpout, spc, nfetch, rA, rB, pA, pB, mu, nu, *offA, *offB, lda, ldb, mulA, mulB, rowA, rowB, ia, ib); } }/* * Drain last of pipe, where all we do is adds */ while (k < lat) { for (j=0; j < nu && k < lat; j++) { for (i=0; i < mu && k < lat; i++, k++) { fprintf(fpout, "%s %s%d_%d += m%d;\n", spc, rC, i, j, h); if (++h == lat) h = 0; } } }}void fetchC(FILE*,char*, enum ATLAS_LOOP_ORDER, int, int, int, int, char*, char*, int, int, int, int, int, char*);/* * regKloop : Creates a full K-loop, unrolled by ku, assuming M & N * unrollings of mu & nu, using separate muladd instruction */void regKloop(FILE *fpout, char *spc, /* indentation string */ enum ATLAS_LOOP_ORDER LoopOrder, enum ATLAS_TRANS TA, enum ATLAS_TRANS TB, int AsgC1, /* if 0, 1st iter does cij = not cij += */ 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, /* if 0, K is arbitrary, else K is len of K-loop */ int ifetch, /* number of initial fetches to perform */ int nfetch, /* # of fetches to perform for every flop */ int lat, /* latency */ char *rA, /* varnam for registers of inner matrix */ char *rB, /* varnam for registers of outer matrix */ char *rC, /* varnam for registers holding C */ char *pA, /* varnam for pointer(s) to inner matrix */ char *pB, /* varnam for pointer(s) to outer matrix */ int mu, /* unrolling along inner loop */ int nu, /* unrolling along outer loop */ int ku, /* unrolling along k-loop (innermost) */ int lda, /* row stride: 0 = unknown */ int ldb, /* row stride: 0 = unknown */ int mulA, /* column stride: 1 = real, 2 = cplx */ int mulB, /* column stride: 1 = real, 2 = cplx */ int incpA, /* Increment A for every K iteration? */ int incpB, /* Increment B for every K iteration? */ char *incAk, /* if !rowA, k-loop increment for ptrs */ char *incBk) /* if !rowB, k-loop increment for ptrs */{ int k, Kb, Kpipe, Kloop, kr; int i, j, h; int rowA, rowB; /* mu/nu elts fetched from row? */ int offA=0, offB=0, offA0, offB0; int ia=0, ib=0, iop=0, jop=0; char incAk0[64], incBk0[64]; sprintf(incAk0, "%s0", incAk); sprintf(incBk0, "%s0", incBk); if (K) { Kb = K - K % ku; if (2*lat > K) { regKloop_ma(fpout, spc, LoopOrder, TA, TB, AsgC1, M, N, K, ifetch, nfetch, lat, rA, rB, rC, pA, pB, mu, nu, ku, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk); return; } assert (ku*2 <= K || K == ku); /* need at least one iter. for loop */ } if (K != ku) { if (K) i = mu*nu*ku; else i = mu*nu; if (i > lat) assert(i == ((i)/lat)*lat); else assert(lat == (lat/i)*i); } if (TA == AtlasNoTrans) rowA = 0; else rowA = 1; if (TB == AtlasNoTrans) rowB = 1; else rowB = 0; if (K == ku) /* fully unrolled loop */ { regKunroll(fpout, spc, LoopOrder, ifetch, nfetch, lat, 1, AsgC1,
⌨️ 快捷键说明
复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?