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