| 4 | using namespace megcc::KernelGen; |
| 5 | |
| 6 | TEST(AARCH64, Int8MatMulM8N12K4Dot) { |
| 7 | Checker<MatrixMulForward> checker(Arch::ARM64); |
| 8 | MatrixMulForward::Param param; |
| 9 | UniformIntRNG rng(-127, 127); |
| 10 | checker.set_rng(0, &rng); |
| 11 | checker.set_rng(1, &rng); |
| 12 | |
| 13 | checker.set_dtype(0, dtype::Int8()); |
| 14 | checker.set_dtype(1, dtype::Int8()); |
| 15 | checker.set_dtype(2, dtype::Int32()); |
| 16 | checker.set_kernel_symbol("Arm64_kernel_int8_dot_matmul_8x12mk4_.*"); |
| 17 | for (size_t m : {4, 8, 16, 64}) |
| 18 | for (size_t n : {3, 8, 15, 56}) |
| 19 | for (size_t k : {4, 8, 16, 64}) { |
| 20 | param.transposeA = false; |
| 21 | param.transposeB = false; |
| 22 | param.format = param::MatrixMul::Format::MK4_DOT; |
| 23 | checker.set_param(param); |
| 24 | checker.execs({{m / 4, k / 4, 4, 4}, {k / 4, n, 4}, {}}); |
| 25 | } |
| 26 | } |
| 27 | |
| 28 | TEST(AARCH64, Int8MatMulM8N12K8MK4I8mm) { |
| 29 | Checker<MatrixMulForward> checker(Arch::ARM64_WITH_I8MM); |