| 58 | #endif |
| 59 | |
| 60 | void sgemv_naive_m( |
| 61 | const float* __restrict A, const float* __restrict B, float* __restrict C, |
| 62 | size_t M, size_t N, size_t K, size_t Astride, size_t Bstride, size_t Cstride) { |
| 63 | size_t m = 0; |
| 64 | for (; m + 4 <= M; m += 4) { |
| 65 | size_t k = 0; |
| 66 | memset(C + m * Cstride, 0, 4 * sizeof(float) * N); |
| 67 | for (; k + 4 <= K; k += 4) { |
| 68 | size_t n = 0; |
| 69 | for (; n + 4 <= N; n += 4) { |
| 70 | float32x4_t a00, a01, a02, a03, a10, a11, a12, a13, a20, a21, a22, a23, |
| 71 | a30, a31, a32, a33; |
| 72 | float32x4_t b0, b1, b2, b3; |
| 73 | float32x4_t c0, c1, c2, c3; |
| 74 | #define loadB(i) b##i = vld1q_f32(B + (k + i) * Bstride + n); |
| 75 | #define loadC(i) c##i = vld1q_f32(C + (m + i) * Cstride + n); |
| 76 | #define loadA0(i) a0##i = vdupq_n_f32(A[(m + 0) * Astride + k + i]); |
| 77 | #define loadA1(i) a1##i = vdupq_n_f32(A[(m + 1) * Astride + k + i]); |
| 78 | #define loadA2(i) a2##i = vdupq_n_f32(A[(m + 2) * Astride + k + i]); |
| 79 | #define loadA3(i) a3##i = vdupq_n_f32(A[(m + 3) * Astride + k + i]); |
| 80 | UNROLL_OUT(loadC, 4) |
| 81 | UNROLL_OUT(loadB, 4) |
| 82 | UNROLL_OUT(loadA0, 4) |
| 83 | UNROLL_OUT(loadA1, 4) |
| 84 | UNROLL_OUT(loadA2, 4) |
| 85 | UNROLL_OUT(loadA3, 4) |
| 86 | #undef loadB |
| 87 | #undef loadC |
| 88 | #undef loadA0 |
| 89 | #undef loadA1 |
| 90 | #undef loadA2 |
| 91 | #undef loadA3 |
| 92 | #define calculate_row0(i) c0 = vmlaq_f32(c0, b##i, a0##i); |
| 93 | #define calculate_row1(i) c1 = vmlaq_f32(c1, b##i, a1##i); |
| 94 | #define calculate_row2(i) c2 = vmlaq_f32(c2, b##i, a2##i); |
| 95 | #define calculate_row3(i) c3 = vmlaq_f32(c3, b##i, a3##i); |
| 96 | UNROLL_OUT(calculate_row0, 4) |
| 97 | UNROLL_OUT(calculate_row1, 4) |
| 98 | UNROLL_OUT(calculate_row2, 4) |
| 99 | UNROLL_OUT(calculate_row3, 4) |
| 100 | #undef calculate_row0 |
| 101 | #undef calculate_row1 |
| 102 | #undef calculate_row2 |
| 103 | #undef calculate_row3 |
| 104 | #define vstore(i) vst1q_f32(C + (m + i) * Cstride + n, c##i); |
| 105 | UNROLL_OUT(vstore, 4) |
| 106 | #undef vstore |
| 107 | } |
| 108 | for (; n + 2 <= N; n += 2) { |
| 109 | float32x4_t a0, a1, a2, a3; |
| 110 | float32x2_t b0, b1, b2, b3; |
| 111 | float32x2_t c0, c1, c2, c3; |
| 112 | #define loadA(i) a##i = vld1q_f32(A + (m + i) * Astride + k); |
| 113 | #define loadB(i) b##i = vld1_f32(B + (k + i) * Bstride + n); |
| 114 | #define loadC(i) c##i = vld1_f32(C + (m + i) * Cstride + n); |
| 115 | UNROLL_OUT(loadC, 4) |
| 116 | UNROLL_OUT(loadA, 4) |
| 117 | UNROLL_OUT(loadB, 4) |