matmul.c

Matrix multiplication helper library
git clone git://git.finwo.net/lib/matmul.c
Log | Files | Refs | README | LICENSE

commit 2ac74e34bb55a7f60b8a1ee4e9d68ce93924efa8
parent 5e4fc7d0e0d87092d880d64949afe6dd13c68bd6
Author: finwo <finwo@pm.me>
Date:   Tue, 21 Apr 2026 20:12:03 +0200

Fix generic generation

Diffstat:
Msrc/matmul.h | 72++++++++++++++++++++++++++----------------------------------------------
1 file changed, 26 insertions(+), 46 deletions(-)

diff --git a/src/matmul.h b/src/matmul.h @@ -48,53 +48,33 @@ extern "C" { extern int (*matmul_u8_i8_u8)(size_t, size_t, size_t, const uint8_t *, const int8_t *, uint8_t *, double); extern int (*matmul_f32_f32_f32)(size_t, size_t, size_t, const float *, const float *, float *, double); extern int (*matmul_f64_f64_f64)(size_t, size_t, size_t, const double *, const double *, double *, double); -extern int (*matmul_not_implemented)(size_t, size_t, size_t, const void *, const void *, void *, double); -#define matmul(m, n, p, A, B, C, scale) \ - _Generic((A), \ - const uint8_t *: _Generic((B), \ - const int8_t *: _Generic((C), uint8_t *: matmul_u8_i8_u8), \ - int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const float *: _Generic((C), uint8_t *: matmul_not_implemented), \ - float *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const double *: _Generic((C), uint8_t *: matmul_not_implemented), \ - double *: _Generic((C), uint8_t *: matmul_not_implemented)), \ - uint8_t *: _Generic((B), \ - const int8_t *: _Generic((C), uint8_t *: matmul_u8_i8_u8), \ - int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const float *: _Generic((C), uint8_t *: matmul_not_implemented), \ - float *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const double *: _Generic((C), uint8_t *: matmul_not_implemented), \ - double *: _Generic((C), uint8_t *: matmul_not_implemented)), \ - const float *: _Generic((B), \ - const int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const float *: _Generic((C), float *: matmul_f32_f32_f32), \ - float *: _Generic((C), float *: matmul_not_implemented), \ - const double *: _Generic((C), float *: matmul_not_implemented), \ - double *: _Generic((C), float *: matmul_not_implemented)), \ - float *: _Generic((B), \ - const int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const float *: _Generic((C), float *: matmul_f32_f32_f32), \ - float *: _Generic((C), float *: matmul_not_implemented), \ - const double *: _Generic((C), float *: matmul_not_implemented), \ - double *: _Generic((C), float *: matmul_not_implemented)), \ - const double *: _Generic((B), \ - const int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const float *: _Generic((C), float *: matmul_not_implemented), \ - float *: _Generic((C), float *: matmul_not_implemented), \ - const double *: _Generic((C), double *: matmul_f64_f64_f64), \ - double *: _Generic((C), double *: matmul_not_implemented)), \ - double *: _Generic((B), \ - const int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - int8_t *: _Generic((C), uint8_t *: matmul_not_implemented), \ - const float *: _Generic((C), float *: matmul_not_implemented), \ - float *: _Generic((C), float *: matmul_not_implemented), \ - const double *: _Generic((C), double *: matmul_f64_f64_f64), \ - double *: _Generic((C), double *: matmul_not_implemented)), \ - void *: matmul_not_implemented)((m), (n), (p), (A), (B), (C), (scale)) +#define __matmul_C(type_a,type_b) \ + _Generic((C), \ + uint8_t *: matmul_##type_a##_##type_b##_u8, \ + float *: matmul_##type_a##_##type_b##_f32, \ + double *: matmul_##type_a##_##type_b##_f64 \ + ) + +#define __matmul_B(type_a) \ + _Generic((B), \ + int8_t *: __matmul_C(type_a,i8) \ + const int8_t *: __matmul_C(type_a,i8) \ + float *: __matmul_C(type_a,f32) \ + const float *: __matmul_C(type_a,f32) \ + double *: __matmul_C(type_a,f64) \ + const double *: __matmul_C(type_a,f64) \ + ) + +#define matmul(m,n,p,A,B,C,scale) \ + _Generic((A), \ + uint8_t *: __matmul_B(u8) \ + const uint8_t *: __matmul_B(u8) \ + float *: __matmul_B(f32) \ + const float *: __matmul_B(f32) \ + double *: __matmul_B(f64) \ + const double *: __matmul_B(f64) \ + ) #ifdef __cplusplus }