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