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