emit_mm.c
来自「基于Blas CLapck的.用过的人知道是干啥的」· C语言 代码 · 共 1,670 行 · 第 1/5 页
C
1,670 行
if (offA) fprintf(fpout, "%s %s0 = %s0[%d];\n", spc, rA, pA, offA); else fprintf(fpout, "%s %s0 = *%s0;\n", spc, rA, pA); if (offB) fprintf(fpout, "%s %s0 = %s0[%d];\n", spc, rB, pB, offB); else fprintf(fpout, "%s %s0 = *%s0;\n", spc, rB, pB); nf = 2; ia = ib = 1; } while ( (nf < ifetch) && (ia < mu || ib < nu) ) { if (ia < mu) /* remaining elts of inner matrix to be fetched */ { if (rowA) /* fetching from row */ { if (lda) fprintf(fpout, "%s %s%d = %s0[%d];\n", spc, rA, ia, pA, offA+(ia*lda)*mulA); else if (offA == 0) fprintf(fpout, "%s %s%d = *%s%d;\n", spc, rA, ia, pA, ia); else fprintf(fpout, "%s %s%d = %s%d[%d];\n", spc, rA, ia, pA, ia, offA); } else fprintf(fpout, "%s %s%d = %s0[%d];\n", spc, rA, ia, pA, offA+ia*mulA); ia++; } else /* after inner matrix fetched, fetch outer matrix */ { if (rowB) /* fetching from row */ { if (ldb) fprintf(fpout, "%s %s%d = %s0[%d];\n", spc, rB, ib, pB, offB+ib*ldb*mulB); else if (offB == 0) fprintf(fpout, "%s %s%d = *%s%d;\n", spc, rB, ib, pB, ib); else fprintf(fpout, "%s %s%d = %s%d[%d];\n", spc, rB, ib, pB, ib, offB); } else fprintf(fpout, "%s %s%d = %s0[%d];\n", spc, rB, ib, pB, offB+ib*mulB); ib++; } nf++; } *ia0 = ia; *ib0 = ib;}/* * incABk : increment A & B pointers or offsets inside K-loop */void incABk(FILE *fpout, char *spc, char *pA, char *pB, /* varnams of pointers to matrices */ int mu, int nu, /* unrollings */ int *offA, int *offB, /* offsets from ptr to first part of block */ int lda, int ldb, /* leading dimensions */ int mulA, int mulB, /* col increment: 1 = real, 2 = cplx */ int incpA, int incpB, /* Increment ptrs at each K iteration? */ char *incAk, char *incBk, /* if needed, K-loop inc constant */ int rowA, int rowB) /* 0: fetch from col, else fetch from row */{ int p; if (incpA) { if (rowA) { if (mulA == 1) { if (lda) fprintf(fpout, "%s %s0++;\n", spc, pA); else for(p=0; p < mu; p++) fprintf(fpout, "%s %s%d++;\n",spc, pA, p); } else { if (lda) fprintf(fpout, "%s %s0 += %d;\n", spc, pA, mulA); else for(p=0; p < mu; p++) fprintf(fpout, "%s %s%d += %d;\n",spc, pA, p, mulA); } } else /* mu unroll is along a column */ { if (lda) fprintf(fpout, "%s %s0 += %d;\n", spc, pA, lda*mulA); else fprintf(fpout, "%s %s0 += %s;\n", spc, pA, incAk); } } else { if (rowA) (*offA) += mulA; else { assert(lda); *offA += mulA*lda; } } if (incpB) { if (rowB) { if (mulB == 1) { if (ldb) fprintf(fpout, "%s %s0++;\n", spc, pB); else for(p=0; p < nu; p++) fprintf(fpout, "%s %s%d++;\n",spc, pB, p); } else { if (ldb) fprintf(fpout, "%s %s0 += %d;\n", spc, pB, mulB); else for(p=0; p < nu; p++) fprintf(fpout, "%s %s%d += %d;\n",spc, pB, p, mulB); } } else /* nu unroll is along a row */ { if (ldb) fprintf(fpout, "%s %s0 += %d;\n", spc, pB, ldb*mulB); else fprintf(fpout, "%s %s0 += %s;\n", spc, pB, incBk); } } else /* incrementing offset, not pointers */ { if (rowB) (*offB) += mulB; else { assert(ldb); *offB += mulB*ldb; } }}/* * regKunroll_ma: K-loop unrolling for combine multiply/add instruction. * 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. */void regKunroll_ma(FILE *fpout, char *spc, /* indentation string */ enum ATLAS_LOOP_ORDER LoopOrder, int Asg1stC,/* 1: first C update gets =, instead of += */ int ifetch, /* number of initial fetches to perform */ int nfetch, /* # of fetches to perform for every flop */ 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 *offA0, /* offset to first elt of this block */ int *offB0, /* 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 i, j, k, ia=0, ib=0, offA=(*offA0), offB=(*offB0); if (LoopOrder == AtlasIJK) { regKunroll_ma(fpout, spc, AtlasJIK, Asg1stC, ifetch, nfetch, rB, rA, rC, pB, pA, nu, mu, ku, offB0, offA0, ldb, lda, mulB, mulA, incpA, incpB, incBk, incAk, rowB, rowA); return; } for (k=0; k < ku; k++) { ia = ib = 0; opfetch(fpout, spc, ifetch, rA, rB, pA, pB, mu, nu, offA, offB, lda, ldb, mulA, mulB, rowA, rowB, &ia, &ib); for (j=0; j < nu; j++) { for (i=0; i < mu; i++) { if (Asg1stC && !k) fprintf(fpout, "%s %s%d_%d = %s%d * %s%d;\n", spc, rC, i, j, rA, i, rB, j); else fprintf(fpout, "%s %s%d_%d += %s%d * %s%d;\n", spc, rC, i, j, rA, i, rB, j); opfetch(fpout, spc, nfetch, rA, rB, pA, pB, mu, nu, offA, offB, lda, ldb, mulA, mulB, rowA, rowB, &ia, &ib); } } incABk(fpout, spc, pA, pB, mu, nu, &offA, &offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); }}/* * regKloop_ma : Creates a full K-loop, unrolled by ku, assuming M & N * unrollings of mu & nu, using combined muladd instruction */void regKloop_ma(FILE *fpout, char *spc, /* indentation string */ enum ATLAS_LOOP_ORDER LoopOrder, enum ATLAS_TRANS TA, enum ATLAS_TRANS TB, int AsgC1, /* if 1, 1st iter does c = instead of c += */ 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 */{ char kadj[8]; int i; int rowA, rowB; /* mu/nu elts fetched from row? */ int offA=0, offB=0; char *incAk0, *incBk0; if (K) assert (ku*2 <= K || K == ku); /* need at least one iter. for loop */ if (TA == AtlasNoTrans) rowA = 0; else rowA = 1; if (TB == AtlasNoTrans) rowB = 1; else rowB = 0; if (ku == K) /* loop fully unrolled */ { regKunroll_ma(fpout, spc, LoopOrder, AsgC1, ifetch, nfetch, rA, rB, rC, pA, pB, mu, nu, ku, &offA, &offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); if (!incpA) { fprintf(fpout, "%s %s0 += %s;\n", spc, pA, incAk); if (lda == 0 && TA != AtlasNoTrans) for (i=1; i < mu; i++) fprintf(fpout, "%s %s%d += %s;\n", spc, pA, i, incAk); } if (!incpB) { fprintf(fpout, "%s %s0 += %s;\n", spc, pB, incBk); if (ldb == 0 && TB == AtlasNoTrans) for (i=1; i < nu; i++) fprintf(fpout, "%s %s%d += %s;\n", spc, pB, i, incBk); } return; } if (K) {/* * If need to do C = on 1st iteration, peal 1st ku iterations, and then * do loop with K-ku */ if (AsgC1 && (K > ku) && ku <= MAX_CASG_KU) { fprintf(fpout, "/*\n *%s Peel 1st iter to assign C regs\n */\n", spc); regKunroll_ma(fpout, spc, LoopOrder, 1, ifetch, nfetch, rA, rB, rC, pA, pB, mu, nu, ku, &offA, &offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); if (!incpA) { fprintf(fpout, "%s %s0 += %s;\n", spc, pA, incAk); if (lda == 0 && TA != AtlasNoTrans) for (i=1; i < mu; i++) fprintf(fpout, "%s %s%d += %s;\n", spc, pA, i, incAk); } if (!incpB) { fprintf(fpout, "%s %s0 += %s;\n", spc, pB, incBk); if (ldb == 0 && TB == AtlasNoTrans) for (i=1; i < nu; i++) fprintf(fpout, "%s %s%d += %s;\n", spc, pB, i, incBk); } fprintf(fpout, "/*\n *%s Unpeeled K iterations\n */\n", spc); K -= ku; }#ifdef ICC_IS_RETARDED if (ku == 1) fprintf(fpout, "%s for (k=0; k < %d; k++) /* easy loop to unroll */\n", spc, K); else fprintf(fpout, "%s for (k=0; k < %d; k++) /* easy loop to unroll */\n", spc, K/ku);#else if (ku == 1) fprintf(fpout, "%s for (k=%d; k; k--) /* easy loop to unroll */\n", spc, K); else fprintf(fpout, "%s for (k=%d; k; k -= %d) /* easy loop to unroll */\n", spc, (K/ku)*ku, ku);#endif } else {#ifdef ICC_IS_RETARDED if (ku == 1) fprintf(fpout, "%s for (k=0; k < K; k++) /* easy loop to unroll */\n", spc); else fprintf(fpout, "%s for (k=0; k < Kb; k += %d) /* easy loop to unroll */\n", spc, ku);#else if (ku == 1) fprintf(fpout, "%s for (k=K; k; k--) /* easy loop to unroll */\n", spc); else fprintf(fpout, "%s for (k=Kb; k; k -= %d) /* easy loop to unroll */\n", spc, ku);#endif } fprintf(fpout, "%s {\n", spc); spc -= 3; regKunroll_ma(fpout, spc, LoopOrder, 0, ifetch, nfetch, rA, rB, rC, pA, pB, mu, nu, ku, &offA, &offB, lda, ldb, mulA, mulB, incpA, incpB, incAk, incBk, rowA, rowB); if (!incpA) /* need to increment the A pointers ourselves */ { fprintf(fpout, "%s %s0 += %s;\n", spc, pA, incAk); if (rowA && !lda)
⌨️ 快捷键说明
复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?