| 4 | using namespace megcc::KernelGen; |
| 5 | #if ENABLE_KERNEL_FP16 |
| 6 | TEST(AARCH64, Fp16MatMulM8N8K8) { |
| 7 | Checker<MatrixMulForward> checker(Arch::ARM64); |
| 8 | MatrixMulForward::Param param; |
| 9 | megcc::test::Float16PeriodicalRNG rng(0x3c00); |
| 10 | // megcc::test::SequenceRNG rng; |
| 11 | checker.set_rng(0, &rng); |
| 12 | checker.set_rng(1, &rng); |
| 13 | checker.set_epsilon(5e-3); |
| 14 | checker.set_dtype(0, dtype::Float16()) |
| 15 | .set_dtype(1, dtype::Float16()) |
| 16 | .set_dtype(2, dtype::Float16()); |
| 17 | |
| 18 | checker.set_kernel_symbol("Arm64_kernel_fp16_matmul_8x8mk8_.*"); |
| 19 | for (size_t m : {8, 16, 64}) |
| 20 | for (size_t n : {3, 8, 15, 56}) |
| 21 | for (size_t k : {8, 16, 64}) { |
| 22 | param.transposeA = false; |
| 23 | param.transposeB = false; |
| 24 | param.format = param::MatrixMul::Format::MK8; |
| 25 | checker.set_param(param); |
| 26 | checker.execs({{m / 8, k / 8, 8, 8}, {k / 8, n, 8}, {}}); |
| 27 | } |
| 28 | } |
| 29 | #endif |
| 30 | // vim: syntax=cpp.doxygen |