atl_gemvn_4x2_0.c

来自「基于Blas CLapck的.用过的人知道是干啥的」· C语言 代码 · 共 646 行 · 第 1/2 页

C
646
字号
         #ifdef BETA1            for (j=(N>>2); j; j--, A += incA, X += 4)               gemvN32x4(M, 4, A, lda, X, ATL_rone, Y);            if ( (j = N-((N>>2)<<2)) ) gemvNle4(M, j, A, lda, X, ATL_rone, Y);         #else            gemvN32x4(M, 4, A, lda, X, beta, Y);            if (N != 4)               Mjoin(PATL,gemvN_a1_x1_b1_y1)                  (M, N-4, ATL_rone, A+incA, lda, X+4, 1, ATL_rone, Y, 1);         #endif      }      else gemvMlt8(M, N, A, lda, X, beta, Y);   }   else if (M) gemvNle4(M, N, A, lda, X, beta, Y);}static void gemvMlt8(const int M, const int N, const TYPE *A, const int lda,                     const TYPE *X, const SCALAR beta, TYPE *Y){   int i;   register TYPE y0;   for (i=M; i; i--)   {      #ifdef BETA0         y0 = Mjoin(PATL,dot)(N, A, lda, X, 1);      #else         Yget(y0, *Y, beta);         y0 += Mjoin(PATL,dot)(N, A, lda, X, 1);      #endif      *Y++ = y0;      A++;   }}static void gemvNle4(const int M, const int N, const TYPE *A, const int lda,                     const TYPE *X, const SCALAR beta, TYPE *Y){   int i;   const TYPE *A0 = A, *A1 = A+lda, *A2 = A1+lda, *A3 = A2+lda;   register TYPE x0, x1, x2, x3;   #ifdef BETAX      const register TYPE bet=beta;   #endif   switch(N)   {   case 1:      #if defined(BETA0)         Mjoin(PATL,cpsc)(M, *X, A, 1, Y, 1);      #elif defined(BETAX)         Mjoin(PATL,axpby)(M, *X, A, 1, beta, Y, 1);      #else         Mjoin(PATL,axpy)(M, *X, A, 1, Y, 1);      #endif      break;   case 2:      x0 = *X; x1 = X[1];      for (i=0; i != M; i++)      #ifdef BETA0         Y[i] = A0[i] * x0 + A1[i] * x1;      #elif defined(BETAX)         Y[i] = Y[i]*bet + A0[i] * x0 + A1[i] * x1;      #else         Y[i] += A0[i] * x0 + A1[i] * x1;      #endif      break;   case 3:      x0 = *X; x1 = X[1]; x2 = X[2];      for (i=0; i != M; i++)      #ifdef BETA0         Y[i] = A0[i] * x0 + A1[i] * x1 + A2[i] * x2;      #elif defined(BETAX)         Y[i] = Y[i]*bet + A0[i] * x0 + A1[i] * x1 + A2[i] * x2;      #else         Y[i] += A0[i] * x0 + A1[i] * x1 + A2[i] * x2;      #endif      break;   case 4:      if (M >= 32) gemv32x4(M, 4, A, lda, X, beta, Y);      else      {         x0 = *X; x1 = X[1]; x2 = X[2]; x3 = X[3];         for (i=0; i != M; i++)         #ifdef BETA0            Y[i] = A0[i] * x0 + A1[i] * x1 + A2[i] * x2 + A3[i] * x3;         #elif defined(BETAX)            Y[i] = Y[i]*bet + A0[i] * x0 + A1[i] * x1 + A2[i] * x2 + A3[i] * x3;         #else            Y[i] += A0[i] * x0 + A1[i] * x1 + A2[i] * x2 + A3[i] * x3;         #endif      }      break;   default:      ATL_assert(!N);   }}void Mjoin(Mjoin(Mjoin(Mjoin(Mjoin(PATL,gemvN),NM),_x1),BNM),_y1)   (const int M, const int N, const SCALAR alpha, const TYPE *A, const int lda,    const TYPE *X, const int incX, const SCALAR beta, TYPE *Y, const int incY){   const int incA = lda<<1, incAm = 4 - ((N>>1)<<1)*lda;   const int m4 = (M>>2)<<2;   int n2, nr;   register TYPE y0, y1, y2, y3, z0, z1, z2, z3, x0, x1, m0, m1, m2, m3;   register TYPE a00, a10, a20, a30, a01, a11, a21, a31;   const TYPE *x, *stX = X + ((N>>1)<<1)-2, *A0 = A, *A1 = A + lda;   TYPE *stY = Y + m4;   if (N > 4)   {      n2 = ((N-4)>>1)<<1;      nr = N - n2;      if (m4)      {         do         {            x = X + 2;            #ifdef BETA0               z0 = z1 = z2 = z3 = y0 = y1 = y2 = y3 = ATL_rzero;            #else               z0 = *Y;               z1 = Y[1];               z2 = Y[2];               z3 = Y[3];               #ifdef BETAX                  y0 = beta;                  z0 *= y0;                  z1 *= y0;                  z2 *= y0;                  z3 *= y0;               #endif               y0 = y1 = y2 = y3 = ATL_rzero;            #endif            x0 = *X;            x1 = X[1];            a00 = *A0;            a01 = *A1;            a10 = A0[1];            a11 = A1[1];            a20 = A0[2];            a21 = A1[2];            a30 = A0[3];            a31 = A1[3];            A0 += incA;            A1 += incA;            m0 = x0 * a00;            a00 = *A0;            m1 = x1 * a01;            a01 = *A1;            m2 = x0 * a10;            a10 = A0[1];            m3 = x1 * a11;            a11 = A1[1];            if (n2)            {               do               {                  y0 += m0;                  m0 = x0 * a20;                  a20 = A0[2];                  z0 += m1;                  m1 = x1 * a21;                  a21 = A1[2];                  y1 += m2;                  m2 = x0 * a30;                  x0 = *x;                  a30 = A0[3];                  A0 += incA;                  z1 += m3;                  m3 = x1 * a31;                  x1 = x[1];                  a31 = A1[3];                  x += 2;                  A1 += incA;                  y2 += m0;                  m0 = x0 * a00;                  a00 = *A0;                  z2 += m1;                  m1 = x1 * a01;                  a01 = *A1;                  y3 += m2;                  m2 = x0 * a10;                  a10 = A0[1];                  z3 += m3;                  m3 = x1 * a11;                  a11 = A1[1];               }               while (x != stX);            }            if (nr == 4)            {               y0 += m0;               m0 = x0 * a20;               a20 = A0[2];               z0 += m1;               m1 = x1 * a21;               a21 = A1[2];               y1 += m2;               m2 = x0 * a30;               x0 = *x;               a30 = A0[3];               z1 += m3;               m3 = x1 * a31;               x1 = x[1];               a31 = A1[3];               y2 += m0;               m0 = x0 * a00;               z2 += m1;               m1 = x1 * a01;               y3 += m2;               m2 = x0 * a10;               z3 += m3;               m3 = x1 * a11;               y0 += m0;               m0 = x0 * a20;               z0 += m1;               m1 = x1 * a21;               y1 += m2;               m2 = x0 * a30;               z1 += m3;               m3 = x1 * a31;               y2 += m0;               A0 += incA;               z2 += m1;               A1 += incA;               y3 += m2;               z3 += m3;            }            else /* nr == 5 */            {               y0 += m0;               m0 = x0 * a20;               a20 = A0[2];               z0 += m1;               m1 = x1 * a21;               a21 = A1[2];               y1 += m2;               m2 = x0 * a30;               x0 = *x;               a30 = A0[3];               A0 += incA;               z1 += m3;               m3 = x1 * a31;               x1 = x[1];               x += 2;               a31 = A1[3];               y2 += m0;               m0 = x0 * a00;               a00 = *A0;               z2 += m1;               m1 = x1 * a01;               y3 += m2;               m2 = x0 * a10;               a10 = A0[1];               z3 += m3;               m3 = x1 * a11;               y0 += m0;               m0 = x0 * a20;               a20 = A0[2];               z0 += m1;               m1 = x1 * a21;               y1 += m2;               m2 = x0 * a30;               x0 = *x;               a30 = A0[3];               z1 += m3;               m3 = x1 * a31;               y2 += m0;               m0 = x0 * a00;               z2 += m1;               m1 = x0 * a10;               y3 += m2;               m2 = x0 * a20;               z3 += m3;               m3 = x0 * a30;               y0 += m0;               A1 += incA;               y1 += m1;               y2 += m2;               y3 += m3;            }            y0 += z0;            A0 += incAm;            y1 += z1;            A1 += incAm;            y2 += z2;            y3 += z3;            *Y = y0;            Y[1] = y1;            Y[2] = y2;            Y[3] = y3;            Y += 4;         }         while (Y != stY);      }      for (nr=M-m4; nr; nr--)      {         #ifdef BETA0            y0 = Mjoin(PATL,dot)(N, A0, lda, X, 1);         #else            #if defined(BETAX)               y0 = *Y * beta;            #else               y0 = *Y;            #endif            y0 += Mjoin(PATL,dot)(N, A0, lda, X, 1);         #endif         *Y++ = y0;         A0++;      }   }   else if (M) gemvNle4(M, N, A, lda, X, beta, Y);}

⌨️ 快捷键说明

复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?