|
1 | 1 | #include "skainet_kernels.h" |
| 2 | +#include "skainet_simd.h" |
2 | 3 |
|
3 | 4 | #include <stddef.h> |
4 | 5 | #include <stdint.h> |
|
20 | 21 | * Dequant per element: `(code - 8) * d`. The `- 8` bias centres the |
21 | 22 | * unsigned 4-bit code. Scale `d` is folded once after the block |
22 | 23 | * accumulator (cheaper than broadcasting it across every inner FMA). |
| 24 | + * |
| 25 | + * Loop order: block OUTER, output row INNER — see q8_0_matmul.c for the |
| 26 | + * rationale. The weight is block-major, so for a fixed block consecutive |
| 27 | + * `o` are 18 bytes apart: weight bytes are read sequentially instead of |
| 28 | + * striding output_dim*18 per step. Per output row the blocks still |
| 29 | + * accumulate in ascending order, so results are numerically identical to |
| 30 | + * the previous row-outer form. |
| 31 | + * |
| 32 | + * NEON path (SKAINET_HAVE_NEON): unpack the 16 code bytes into two |
| 33 | + * int8x16 nibble vectors (vand / vshr), re-centre with vsubq_s8(.., 8), |
| 34 | + * widen to f32 lanes and FMA against the activation lanes — the same |
| 35 | + * widen+vfmaq_f32 structure as q8_0_matmul.c, because activations are |
| 36 | + * FP32 (an int8 dotprod would need int8 activations). Plain NEON only, |
| 37 | + * no dotprod/i8mm requirement, so the body runs on every AArch64 core. |
23 | 38 | */ |
24 | 39 |
|
25 | 40 | /* Portable FP16 → FP32 conversion. Matches the Kotlin |
@@ -68,27 +83,65 @@ SKAINET_API void skainet_q4_0_matmul( |
68 | 83 | const int32_t BLOCK_SIZE = 32; |
69 | 84 | const int32_t BYTES_PER_BLOCK = 18; |
70 | 85 | const int32_t blocks_per_input_dim = input_dim / BLOCK_SIZE; |
| 86 | + float* SKAINET_RESTRICT out_base = output + output_offset; |
| 87 | + |
| 88 | + for (int32_t o = 0; o < output_dim; ++o) out_base[o] = 0.0f; |
| 89 | + |
| 90 | + for (int32_t block_idx = 0; block_idx < blocks_per_input_dim; ++block_idx) { |
| 91 | + const float* SKAINET_RESTRICT input_block = |
| 92 | + input + input_offset + (size_t) block_idx * BLOCK_SIZE; |
| 93 | + const uint8_t* SKAINET_RESTRICT block = |
| 94 | + weight + weight_byte_offset + |
| 95 | + (size_t)(block_idx * output_dim) * BYTES_PER_BLOCK; |
71 | 96 |
|
72 | | - for (int32_t o = 0; o < output_dim; ++o) { |
73 | | - float acc = 0.0f; |
74 | | - for (int32_t block_idx = 0; block_idx < blocks_per_input_dim; ++block_idx) { |
75 | | - const uint8_t* SKAINET_RESTRICT block = |
76 | | - weight + weight_byte_offset + |
77 | | - (size_t)(block_idx * output_dim + o) * BYTES_PER_BLOCK; |
| 97 | + for (int32_t o = 0; o < output_dim; ++o, block += BYTES_PER_BLOCK) { |
78 | 98 | uint16_t d_bits = (uint16_t) block[0] | ((uint16_t) block[1] << 8); |
79 | 99 | float d = skainet_q4_0_fp16_to_fp32(d_bits); |
80 | 100 | const uint8_t* SKAINET_RESTRICT codes = block + 2; |
81 | | - const float* SKAINET_RESTRICT input_block = |
82 | | - input + input_offset + (size_t) block_idx * BLOCK_SIZE; |
83 | 101 | float block_sum = 0.0f; |
| 102 | +#ifdef SKAINET_HAVE_NEON |
| 103 | + /* Split-nibble unpack: low nibbles are elements 0..15, high |
| 104 | + * nibbles are elements 16..31. Re-centre by 8 in the signed |
| 105 | + * int8 domain, then widen to f32 and FMA — same structure as |
| 106 | + * the q8_0 NEON body. */ |
| 107 | + const uint8x16_t packed = vld1q_u8(codes); |
| 108 | + const int8x16_t lo8 = vsubq_s8( |
| 109 | + vreinterpretq_s8_u8(vandq_u8(packed, vdupq_n_u8(0x0F))), |
| 110 | + vdupq_n_s8(8)); |
| 111 | + const int8x16_t hi8 = vsubq_s8( |
| 112 | + vreinterpretq_s8_u8(vshrq_n_u8(packed, 4)), |
| 113 | + vdupq_n_s8(8)); |
| 114 | + const int16x8_t lo16a = vmovl_s8(vget_low_s8(lo8)); /* elems 0..7 */ |
| 115 | + const int16x8_t lo16b = vmovl_s8(vget_high_s8(lo8)); /* elems 8..15 */ |
| 116 | + const int16x8_t hi16a = vmovl_s8(vget_low_s8(hi8)); /* elems 16..23 */ |
| 117 | + const int16x8_t hi16b = vmovl_s8(vget_high_s8(hi8)); /* elems 24..31 */ |
| 118 | + float32x4_t accv = vdupq_n_f32(0.0f); |
| 119 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 0), |
| 120 | + vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo16a)))); |
| 121 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 4), |
| 122 | + vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo16a)))); |
| 123 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 8), |
| 124 | + vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo16b)))); |
| 125 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 12), |
| 126 | + vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo16b)))); |
| 127 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 16), |
| 128 | + vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi16a)))); |
| 129 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 20), |
| 130 | + vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi16a)))); |
| 131 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 24), |
| 132 | + vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi16b)))); |
| 133 | + accv = vfmaq_f32(accv, vld1q_f32(input_block + 28), |
| 134 | + vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi16b)))); |
| 135 | + block_sum = skainet_neon_hadd_f32(accv); |
| 136 | +#else |
84 | 137 | for (int32_t k = 0; k < 16; ++k) { |
85 | 138 | int32_t lo = (int32_t)(codes[k] & 0x0F) - 8; |
86 | 139 | int32_t hi = (int32_t)(codes[k] >> 4) - 8; |
87 | 140 | block_sum += input_block[k] * (float) lo; |
88 | 141 | block_sum += input_block[k + 16] * (float) hi; |
89 | 142 | } |
90 | | - acc += block_sum * d; |
| 143 | +#endif |
| 144 | + out_base[o] += block_sum * d; |
91 | 145 | } |
92 | | - output[output_offset + o] = acc; |
93 | 146 | } |
94 | 147 | } |
0 commit comments