diff --git a/skainet-backends/skainet-backend-native-cpu/native/src/bf16_matmul.c b/skainet-backends/skainet-backend-native-cpu/native/src/bf16_matmul.c index f8f04c71..923a42c8 100644 --- a/skainet-backends/skainet-backend-native-cpu/native/src/bf16_matmul.c +++ b/skainet-backends/skainet-backend-native-cpu/native/src/bf16_matmul.c @@ -21,12 +21,29 @@ * (BF16 shares the FP32 sign and exponent layout; only the trailing 16 * mantissa bits are discarded). * - * Iteration order is i-p-j (outer-product into rows of C). The inner - * `c[j] += a_ip * bf16_to_float(b[j])` loop streams two contiguous - * arrays — auto-vectorizes under -O3 -ffast-math into vfmadd231ps - * (x86_64) / fmla (AArch64). On ARMv8.6-A+ a future pass can swap the - * scalar dequant for a `bfdot`/`bfmmla` intrinsic kernel; that lives - * behind a runtime feature check and is out of scope here. + * Iteration order depends on m. + * + * At m == 1 it is plain i-p-j (outer-product into the single row of C). The + * inner `c[j] += a_ip * bf16_to_float(b[j])` loop streams two contiguous + * arrays — auto-vectorizes under -O3 -ffast-math into vfmadd231ps (x86_64) / + * fmla (AArch64). Every B element is used exactly once, so there is nothing to + * reuse and sequential streaming is the right shape. + * + * At m > 1, i-p-j walks the whole of B once per row of A. For ffn_up 8B at + * m=16 that is 16 passes over 90 MiB, and the kernel is bandwidth-bound long + * before it is dequant-bound — the shift itself is nearly free. So j is tiled, + * and within a tile each B row is widened once into a small stack buffer and + * multiplied into all m rows of C. B is then read once in total rather than m + * times. The tile is sized so the widened row and the m C rows it feeds stay + * resident together. + * + * Accumulation order into any given C element is p ascending on both paths, so + * the two are bit-identical to each other and to the original i-p-j + * formulation, not merely close. + * + * On ARMv8.6-A+ a future pass can swap the scalar dequant for a + * `bfdot`/`bfmmla` intrinsic kernel; that lives behind a runtime feature check + * and is out of scope here. * * Caller contract (mirrors skainet_fp32_matmul): * - C is FULLY OVERWRITTEN in the m×n block. @@ -34,6 +51,13 @@ * - m == 0 || n == 0 is a no-op. * - Negative m / n / k are caller errors; defensively treated as no-op. */ +/* + * Columns widened per pass. 512 floats is a 2 KiB stack buffer — small enough + * to leave the C rows it feeds resident alongside it, large enough that the + * per-tile loop overhead disappears against the k*m inner work. + */ +#define SKAINET_BF16_TILE 512 + SKAINET_API void skainet_bf16_matmul( const float* SKAINET_RESTRICT a, int32_t a_offset, int32_t a_stride, const uint8_t* SKAINET_RESTRICT b, int32_t b_byte_offset, int32_t b_byte_stride, @@ -52,9 +76,9 @@ SKAINET_API void skainet_bf16_matmul( } if (k <= 0) return; - for (int32_t i = 0; i < m; ++i) { - const float* SKAINET_RESTRICT a_row = a + a_offset + (size_t) i * a_stride; - float* SKAINET_RESTRICT c_row = c + c_offset + (size_t) i * c_stride; + if (m == 1) { + const float* SKAINET_RESTRICT a_row = a + a_offset; + float* SKAINET_RESTRICT c_row = c + c_offset; for (int32_t p = 0; p < k; ++p) { const float a_ip = a_row[p]; const uint8_t* SKAINET_RESTRICT b_row = @@ -70,5 +94,35 @@ SKAINET_API void skainet_bf16_matmul( c_row[j] += a_ip * b_pj; } } + return; + } + + float widened[SKAINET_BF16_TILE]; + + for (int32_t j0 = 0; j0 < n; j0 += SKAINET_BF16_TILE) { + const int32_t tile = + (n - j0) < SKAINET_BF16_TILE ? (n - j0) : SKAINET_BF16_TILE; + + for (int32_t p = 0; p < k; ++p) { + const uint8_t* SKAINET_RESTRICT b_row = + b + b_byte_offset + (size_t) p * b_byte_stride + (size_t) j0 * 2; + + /* Widen this row's tile once, for all m rows of C below. */ + for (int32_t j = 0; j < tile; ++j) { + uint16_t bits; + memcpy(&bits, b_row + (size_t) j * 2, sizeof(uint16_t)); + uint32_t fp32_bits = ((uint32_t) bits) << 16; + memcpy(&widened[j], &fp32_bits, sizeof(float)); + } + + for (int32_t i = 0; i < m; ++i) { + const float a_ip = a[a_offset + (size_t) i * a_stride + p]; + float* SKAINET_RESTRICT c_row = + c + c_offset + (size_t) i * c_stride + j0; + for (int32_t j = 0; j < tile; ++j) { + c_row[j] += a_ip * widened[j]; + } + } + } } } diff --git a/skainet-backends/skainet-backend-native-cpu/src/jvmTest/kotlin/sk/ainet/exec/kernel/NativeBf16MatmulKernelParityTest.kt b/skainet-backends/skainet-backend-native-cpu/src/jvmTest/kotlin/sk/ainet/exec/kernel/NativeBf16MatmulKernelParityTest.kt index 350a8ed2..a4166799 100644 --- a/skainet-backends/skainet-backend-native-cpu/src/jvmTest/kotlin/sk/ainet/exec/kernel/NativeBf16MatmulKernelParityTest.kt +++ b/skainet-backends/skainet-backend-native-cpu/src/jvmTest/kotlin/sk/ainet/exec/kernel/NativeBf16MatmulKernelParityTest.kt @@ -143,6 +143,53 @@ class NativeBf16MatmulKernelParityTest { ) } + @Test + fun multi_tile_n_with_partial_last_tile_matches_panama() { + // At m > 1 the kernel tiles j at 512 columns. Every other shape here is + // either m == 1 or n <= 256, so without this case the tiled path runs + // as a single full tile and the tile-boundary arithmetic is never + // exercised. n = 1100 is two full tiles plus a 76-column remainder. + val rng = Random(7) + val m = 3; val n = 1100; val k = 17 + val a = FloatArray(m * k) { rng.nextFloat() - 0.5f } + val bFloats = FloatArray(k * n) { rng.nextFloat() - 0.5f } + val b = bf16Bytes(bFloats) + assertParity( + m = m, n = n, k = k, + a = a, aOffset = 0, aStride = k, + b = b, bByteOffset = 0, bByteStride = n * 2, + outStride = n, + ) + } + + @Test + fun tiled_and_single_row_paths_agree_on_the_same_weights() { + // m == 1 and m > 1 take different loop orders. Accumulation stays p + // ascending in both, so row 0 of a multi-row call must be bit-identical + // to the same row computed on its own — not merely within tolerance. + val rng = Random(31) + val n = 700; val k = 9 + val bFloats = FloatArray(k * n) { rng.nextFloat() - 0.5f } + val b = bf16Bytes(bFloats) + val aRow = FloatArray(k) { rng.nextFloat() - 0.5f } + + val single = FloatArray(n) + NativeBf16MatmulKernel.matmul(aRow, 0, k, b, 0, n * 2, single, 0, n, 1, n, k) + + val a2 = FloatArray(2 * k) + aRow.copyInto(a2, 0) + aRow.copyInto(a2, k) + val pair = FloatArray(2 * n) + NativeBf16MatmulKernel.matmul(a2, 0, k, b, 0, n * 2, pair, 0, n, 2, n, k) + + for (j in 0 until n) { + assertEquals( + single[j].toRawBits(), pair[j].toRawBits(), + "tiled path diverged from the single-row path at column $j", + ) + } + } + @Test fun zero_m_or_n_no_op() { val out = FloatArray(5) { 7f }