matrix_mult.hpp
来自「矩阵运算源码最新版本」· HPP 代码 · 共 1,560 行 · 第 1/4 页
HPP
1,560 行
d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt0; d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt0; d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt0; d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt0; d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt0; d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt0; d14+=_mm_load_pd(&C[0+(i+14)*stride+2*k])*bt0; } _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+0)], d00); _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+0)], d02); _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+0)], d04); _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+0)], d06); _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+0)], d08); _mm_store_pd(&D[0+(i+10)*stride+2*(j+0)], d10); _mm_store_pd(&D[0+(i+12)*stride+2*(j+0)], d12); _mm_store_pd(&D[0+(i+14)*stride+2*(j+0)], d14); } { __m128d d00 = _mm_load_pd(&D[0+(i+ 0)*stride+2*(j+1)]); __m128d d02 = _mm_load_pd(&D[0+(i+ 2)*stride+2*(j+1)]); __m128d d04 = _mm_load_pd(&D[0+(i+ 4)*stride+2*(j+1)]); __m128d d06 = _mm_load_pd(&D[0+(i+ 6)*stride+2*(j+1)]); __m128d d08 = _mm_load_pd(&D[0+(i+ 8)*stride+2*(j+1)]); __m128d d10 = _mm_load_pd(&D[0+(i+10)*stride+2*(j+1)]); __m128d d12 = _mm_load_pd(&D[0+(i+12)*stride+2*(j+1)]); __m128d d14 = _mm_load_pd(&D[0+(i+14)*stride+2*(j+1)]); for (int k = 0; k < baseOrder; k++) { __m128d bt0 = _mm_load1_pd(&BT[1+j*stride+2*k]); d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt0; d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt0; d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt0; d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt0; d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt0; d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt0; d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt0; d14+=_mm_load_pd(&C[0+(i+14)*stride+2*k])*bt0; } _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+1)], d00); _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+1)], d02); _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+1)], d04); _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+1)], d06); _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+1)], d08); _mm_store_pd(&D[0+(i+10)*stride+2*(j+1)], d10); _mm_store_pd(&D[0+(i+12)*stride+2*(j+1)], d12); _mm_store_pd(&D[0+(i+14)*stride+2*(j+1)], d14); } } #endif #if 0 // Tweak SSE #define MM_LOAD1_PD(a,b) \ { \ __asm__("movlpd %1, %0" : "=x" (a) : "m"(*b)); \ __asm__("movhpd %1, %0" : "=x" (a) : "m"(*b), "0" (a)); \ } #define MM_LOAD1U_PD(a,b) \ { \ __asm__("movlpd %1, %0" : "=x" (a) : "m"(*b)); \ __asm__("movhpd %1, %0" : "=x" (a) : "m"(*b), "0" (a)); \ } #define MM_MUL_PD(out,addr) \ { out = _mm_mul_pd(out, *(__m128d*)addr); } for (int j = 0; j < baseOrder; j+=2) for (int i = 0; i < baseOrder; i+=16) { { __m128d d00 = _mm_load_pd(&D[0+(i+ 0)*stride+2*(j+0)]); __m128d d02 = _mm_load_pd(&D[0+(i+ 2)*stride+2*(j+0)]); __m128d d04 = _mm_load_pd(&D[0+(i+ 4)*stride+2*(j+0)]); __m128d d06 = _mm_load_pd(&D[0+(i+ 6)*stride+2*(j+0)]); __m128d d08 = _mm_load_pd(&D[0+(i+ 8)*stride+2*(j+0)]); __m128d d10 = _mm_load_pd(&D[0+(i+10)*stride+2*(j+0)]); __m128d d12 = _mm_load_pd(&D[0+(i+12)*stride+2*(j+0)]); __m128d d14 = _mm_load_pd(&D[0+(i+14)*stride+2*(j+0)]); for (int k = 0; k < baseOrder; k++) { __m128d bt0; MM_LOAD1_PD(bt0, &BT[0+j*stride+2*k]); d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt0; d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt0; d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt0; d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt0; d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt0; d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt0; d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt0; MM_MUL_PD(bt0, &C[0+(i+14)*stride+2*k]); d14+=bt0; } _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+0)], d00); _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+0)], d02); _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+0)], d04); _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+0)], d06); _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+0)], d08); _mm_store_pd(&D[0+(i+10)*stride+2*(j+0)], d10); _mm_store_pd(&D[0+(i+12)*stride+2*(j+0)], d12); _mm_store_pd(&D[0+(i+14)*stride+2*(j+0)], d14); } { __m128d d00 = _mm_load_pd(&D[0+(i+ 0)*stride+2*(j+1)]); __m128d d02 = _mm_load_pd(&D[0+(i+ 2)*stride+2*(j+1)]); __m128d d04 = _mm_load_pd(&D[0+(i+ 4)*stride+2*(j+1)]); __m128d d06 = _mm_load_pd(&D[0+(i+ 6)*stride+2*(j+1)]); __m128d d08 = _mm_load_pd(&D[0+(i+ 8)*stride+2*(j+1)]); __m128d d10 = _mm_load_pd(&D[0+(i+10)*stride+2*(j+1)]); __m128d d12 = _mm_load_pd(&D[0+(i+12)*stride+2*(j+1)]); __m128d d14 = _mm_load_pd(&D[0+(i+14)*stride+2*(j+1)]); for (int k = 0; k < baseOrder; k++) { __m128d bt1; MM_LOAD1U_PD(bt1, &BT[1+j*stride+2*k]); d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt1; d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt1; d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt1; d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt1; d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt1; d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt1; d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt1; MM_MUL_PD(bt1, &C[0+(i+14)*stride+2*k]); d14+=bt1; } _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+1)], d00); _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+1)], d02); _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+1)], d04); _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+1)], d06); _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+1)], d08); _mm_store_pd(&D[0+(i+10)*stride+2*(j+1)], d10); _mm_store_pd(&D[0+(i+12)*stride+2*(j+1)], d12); _mm_store_pd(&D[0+(i+14)*stride+2*(j+1)], d14); } } #endif #if 0 // Factor and unroll k #define MM_LOAD1_PD(a,b) \ { \ __asm__("movlpd %1, %0" : "=x" (a) : "m"(*b)); \ __asm__("movhpd %1, %0" : "=x" (a) : "m"(*b), "0" (a)); \ } #define MM_LOAD1U_PD(a,b) \ { \ __asm__("movlpd %1, %0" : "=x" (a) : "m"(*b)); \ __asm__("movhpd %1, %0" : "=x" (a) : "m"(*b), "0" (a)); \ } #define MM_MUL_PD(out,addr) \ { out = _mm_mul_pd(out, *(__m128d*)addr); } #define BLOCK0_0(i,j,k) \ { \ __m128d bt0; \ MM_LOAD1_PD(bt0, &BT[0+j*stride+2*k]); \ d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt0; \ d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt0; \ d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt0; \ d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt0; \ d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt0; \ d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt0; \ d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt0; \ MM_MUL_PD(bt0, &C[0+(i+14)*stride+2*k]); \ d14+=bt0; \ } #define BLOCK0_1(i,j,k) \ { \ __m128d bt1; \ MM_LOAD1U_PD(bt1, &BT[1+j*stride+2*k]); \ d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt1; \ d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt1; \ d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt1; \ d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt1; \ d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt1; \ d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt1; \ d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt1; \ MM_MUL_PD(bt1, &C[0+(i+14)*stride+2*k]); \ d14+=bt1; \ } for (int j = 0; j < baseOrder; j+=2) for (int i = 0; i < baseOrder; i+=16) { { __m128d d00 = _mm_load_pd(&D[0+(i+ 0)*stride+2*(j+0)]); __m128d d02 = _mm_load_pd(&D[0+(i+ 2)*stride+2*(j+0)]); __m128d d04 = _mm_load_pd(&D[0+(i+ 4)*stride+2*(j+0)]); __m128d d06 = _mm_load_pd(&D[0+(i+ 6)*stride+2*(j+0)]); __m128d d08 = _mm_load_pd(&D[0+(i+ 8)*stride+2*(j+0)]); __m128d d10 = _mm_load_pd(&D[0+(i+10)*stride+2*(j+0)]); __m128d d12 = _mm_load_pd(&D[0+(i+12)*stride+2*(j+0)]); __m128d d14 = _mm_load_pd(&D[0+(i+14)*stride+2*(j+0)]); for (int k = 0; k < baseOrder; k+=32) { BLOCK0_0(i,j,(k+ 0)); BLOCK0_0(i,j,(k+ 1)); BLOCK0_0(i,j,(k+ 2)); BLOCK0_0(i,j,(k+ 3)); BLOCK0_0(i,j,(k+ 4)); BLOCK0_0(i,j,(k+ 5)); BLOCK0_0(i,j,(k+ 6)); BLOCK0_0(i,j,(k+ 7)); BLOCK0_0(i,j,(k+ 8)); BLOCK0_0(i,j,(k+ 9)); BLOCK0_0(i,j,(k+10)); BLOCK0_0(i,j,(k+11)); BLOCK0_0(i,j,(k+12)); BLOCK0_0(i,j,(k+13)); BLOCK0_0(i,j,(k+14)); BLOCK0_0(i,j,(k+15)); BLOCK0_0(i,j,(k+16)); BLOCK0_0(i,j,(k+17)); BLOCK0_0(i,j,(k+18)); BLOCK0_0(i,j,(k+19)); BLOCK0_0(i,j,(k+20)); BLOCK0_0(i,j,(k+21)); BLOCK0_0(i,j,(k+22)); BLOCK0_0(i,j,(k+23)); BLOCK0_0(i,j,(k+24)); BLOCK0_0(i,j,(k+25)); BLOCK0_0(i,j,(k+26)); BLOCK0_0(i,j,(k+27)); BLOCK0_0(i,j,(k+28)); BLOCK0_0(i,j,(k+29)); BLOCK0_0(i,j,(k+30)); BLOCK0_0(i,j,(k+31)); } _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+0)], d00); _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+0)], d02); _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+0)], d04); _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+0)], d06); _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+0)], d08); _mm_store_pd(&D[0+(i+10)*stride+2*(j+0)], d10); _mm_store_pd(&D[0+(i+12)*stride+2*(j+0)], d12); _mm_store_pd(&D[0+(i+14)*stride+2*(j+0)], d14); } { __m128d d00 = _mm_load_pd(&D[0+(i+ 0)*stride+2*(j+1)]); __m128d d02 = _mm_load_pd(&D[0+(i+ 2)*stride+2*(j+1)]); __m128d d04 = _mm_load_pd(&D[0+(i+ 4)*stride+2*(j+1)]); __m128d d06 = _mm_load_pd(&D[0+(i+ 6)*stride+2*(j+1)]); __m128d d08 = _mm_load_pd(&D[0+(i+ 8)*stride+2*(j+1)]); __m128d d10 = _mm_load_pd(&D[0+(i+10)*stride+2*(j+1)]); __m128d d12 = _mm_load_pd(&D[0+(i+12)*stride+2*(j+1)]); __m128d d14 = _mm_load_pd(&D[0+(i+14)*stride+2*(j+1)]); for (int k = 0; k < baseOrder; k+=32) { BLOCK0_1(i,j,(k+ 0)); BLOCK0_1(i,j,(k+ 1)); BLOCK0_1(i,j,(k+ 2)); BLOCK0_1(i,j,(k+ 3)); BLOCK0_1(i,j,(k+ 4)); BLOCK0_1(i,j,(k+ 5)); BLOCK0_1(i,j,(k+ 6)); BLOCK0_1(i,j,(k+ 7)); BLOCK0_1(i,j,(k+ 8)); BLOCK0_1(i,j,(k+ 9)); BLOCK0_1(i,j,(k+10)); BLOCK0_1(i,j,(k+11)); BLOCK0_1(i,j,(k+12)); BLOCK0_1(i,j,(k+13)); BLOCK0_1(i,j,(k+14)); BLOCK0_1(i,j,(k+15)); BLOCK0_1(i,j,(k+16)); BLOCK0_1(i,j,(k+17)); BLOCK0_1(i,j,(k+18)); BLOCK0_1(i,j,(k+19)); BLOCK0_1(i,j,(k+20)); BLOCK0_1(i,j,(k+21)); BLOCK0_1(i,j,(k+22)); BLOCK0_1(i,j,(k+23)); BLOCK0_1(i,j,(k+24)); BLOCK0_1(i,j,(k+25)); BLOCK0_1(i,j,(k+26)); BLOCK0_1(i,j,(k+27)); BLOCK0_1(i,j,(k+28)); BLOCK0_1(i,j,(k+29)); BLOCK0_1(i,j,(k+30)); BLOCK0_1(i,j,(k+31)); } _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+1)], d00); _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+1)], d02); _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+1)], d04); _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+1)], d06); _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+1)], d08); _mm_store_pd(&D[0+(i+10)*stride+2*(j+1)], d10); _mm_store_pd(&D[0+(i+12)*stride+2*(j+1)], d12); _mm_store_pd(&D[0+(i+14)*stride+2*(j+1)], d14); } } #endif #if 1 // Factor and unroll i #define MM_LOAD1_PD(a,b) \ { \ __asm__("movlpd %1, %0" : "=x" (a) : "m"(*b)); \ __asm__("movhpd %1, %0" : "=x" (a) : "m"(*b), "0" (a)); \ } #define MM_LOAD1U_PD(a,b) \ { \ __asm__("movlpd %1, %0" : "=x" (a) : "m"(*b)); \ __asm__("movhpd %1, %0" : "=x" (a) : "m"(*b), "0" (a)); \ } #define MM_MUL_PD(out,addr) \ { out = _mm_mul_pd(out, *(__m128d*)addr); } #define BLOCK0_0(i,j,k) \ { \ __m128d bt0; \ MM_LOAD1_PD(bt0, &BT[0+j*stride+2*k]); \ d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt0; \ d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt0; \ d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt0; \ d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt0; \ d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt0; \ d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt0; \ d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt0; \ MM_MUL_PD(bt0, &C[0+(i+14)*stride+2*k]); \ d14+=bt0; \ } #define BLOCK0_1(i,j,k) \ { \ __m128d bt1; \ MM_LOAD1U_PD(bt1, &BT[1+j*stride+2*k]); \ d00+=_mm_load_pd(&C[0+(i+ 0)*stride+2*k])*bt1; \ d02+=_mm_load_pd(&C[0+(i+ 2)*stride+2*k])*bt1; \ d04+=_mm_load_pd(&C[0+(i+ 4)*stride+2*k])*bt1; \ d06+=_mm_load_pd(&C[0+(i+ 6)*stride+2*k])*bt1; \ d08+=_mm_load_pd(&C[0+(i+ 8)*stride+2*k])*bt1; \ d10+=_mm_load_pd(&C[0+(i+10)*stride+2*k])*bt1; \ d12+=_mm_load_pd(&C[0+(i+12)*stride+2*k])*bt1; \ MM_MUL_PD(bt1, &C[0+(i+14)*stride+2*k]); \ d14+=bt1; \ } #define BLOCK1_0(i,j) \ { \ __m128d d00 = _mm_load_pd(&D[0+(i+ 0)*stride+2*(j+0)]); \ __m128d d02 = _mm_load_pd(&D[0+(i+ 2)*stride+2*(j+0)]); \ __m128d d04 = _mm_load_pd(&D[0+(i+ 4)*stride+2*(j+0)]); \ __m128d d06 = _mm_load_pd(&D[0+(i+ 6)*stride+2*(j+0)]); \ __m128d d08 = _mm_load_pd(&D[0+(i+ 8)*stride+2*(j+0)]); \ __m128d d10 = _mm_load_pd(&D[0+(i+10)*stride+2*(j+0)]); \ __m128d d12 = _mm_load_pd(&D[0+(i+12)*stride+2*(j+0)]); \ __m128d d14 = _mm_load_pd(&D[0+(i+14)*stride+2*(j+0)]); \ for (int k = 0; k < baseOrder; k+=32) \ { \ BLOCK0_0(i,j,(k+ 0)); \ BLOCK0_0(i,j,(k+ 1)); \ BLOCK0_0(i,j,(k+ 2)); \ BLOCK0_0(i,j,(k+ 3)); \ BLOCK0_0(i,j,(k+ 4)); \ BLOCK0_0(i,j,(k+ 5)); \ BLOCK0_0(i,j,(k+ 6)); \ BLOCK0_0(i,j,(k+ 7)); \ BLOCK0_0(i,j,(k+ 8)); \ BLOCK0_0(i,j,(k+ 9)); \ BLOCK0_0(i,j,(k+10)); \ BLOCK0_0(i,j,(k+11)); \ BLOCK0_0(i,j,(k+12)); \ BLOCK0_0(i,j,(k+13)); \ BLOCK0_0(i,j,(k+14)); \ BLOCK0_0(i,j,(k+15)); \ BLOCK0_0(i,j,(k+16)); \ BLOCK0_0(i,j,(k+17)); \ BLOCK0_0(i,j,(k+18)); \ BLOCK0_0(i,j,(k+19)); \ BLOCK0_0(i,j,(k+20)); \ BLOCK0_0(i,j,(k+21)); \ BLOCK0_0(i,j,(k+22)); \ BLOCK0_0(i,j,(k+23)); \ BLOCK0_0(i,j,(k+24)); \ BLOCK0_0(i,j,(k+25)); \ BLOCK0_0(i,j,(k+26)); \ BLOCK0_0(i,j,(k+27)); \ BLOCK0_0(i,j,(k+28)); \ BLOCK0_0(i,j,(k+29)); \ BLOCK0_0(i,j,(k+30)); \ BLOCK0_0(i,j,(k+31)); \ } \ _mm_store_pd(&D[0+(i+ 0)*stride+2*(j+0)], d00); \ _mm_store_pd(&D[0+(i+ 2)*stride+2*(j+0)], d02); \ _mm_store_pd(&D[0+(i+ 4)*stride+2*(j+0)], d04); \ _mm_store_pd(&D[0+(i+ 6)*stride+2*(j+0)], d06); \ _mm_store_pd(&D[0+(i+ 8)*stride+2*(j+0)], d08); \ _mm_store_pd(&D[0+(i+10)*stride+2*(j+0)], d10); \ _mm_store_pd(&D[0+(i+12)*stride+2*(j+0)], d12); \ _mm_store_pd(&D[0+(i+14)*stride+2*(j+0)], d14); \ } #define BLOCK1_1(i,j) \ { \
⌨️ 快捷键说明
复制代码Ctrl + C
搜索代码Ctrl + F
全屏模式F11
增大字号Ctrl + =
减小字号Ctrl + -
显示快捷键?