From 5474e582cade0aa80e3bb8045ae9430faa51d2f0 Mon Sep 17 00:00:00 2001
From: Qxiang Xu <Qixiang.Xu@arm.com>
Date: Tue, 25 Aug 2026 10:05:42 +0800
Subject: [PATCH 4/4] Dispatch QD8 F16 QC4W through KleidiAI SME2

---
 cmake/gen/neonsme2_microkernels.cmake         |  4 +-
 gen/neonsme2_microkernels.bzl                 |  2 +-
 src/configs/gemm-config.c                     | 25 +++++++++++
 ...d8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c | 41 +++++++++++++++++++
 src/xnnpack/gemm.h                            |  2 +
 5 files changed, 71 insertions(+), 3 deletions(-)

diff --git a/cmake/gen/neonsme2_microkernels.cmake b/cmake/gen/neonsme2_microkernels.cmake
index 45d7de4c3c..66a08721df 100644
--- a/cmake/gen/neonsme2_microkernels.cmake
+++ b/cmake/gen/neonsme2_microkernels.cmake
@@ -19,6 +19,7 @@ SET(PROD_NEONSME2_MICROKERNEL_SRCS
   src/pqs8-f32-qc8w-igemm/pqs8-f32-qc8w-igemm-32x32c4-minmax-neonsme2.c
   src/pqs8-qc8w-gemm/pqs8-qc8w-gemm-1x32c4-minmax-neonsme2.c
   src/pqs8-qc8w-gemm/pqs8-qc8w-gemm-32x32c4-minmax-neonsme2.c
+  src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c
   src/qp8-f32-qc4w-gemm/qp8-f32-qc4w-gemm-minmax-1x64c4-neonsme2.c
   src/qp8-f32-qc4w-gemm/qp8-f32-qc4w-gemm-minmax-16x64c4-neonsme2.c
   src/qp8-f32-qc8w-gemm/qp8-f32-qc8w-gemm-minmax-1x64c4-neonsme2.c
@@ -30,7 +31,6 @@ SET(PROD_NEONSME2_MICROKERNEL_SRCS
   src/x32-pack-lh/x32-packlh-igemm-neonsme2.c
   src/x32-pack-lh/x32-packlh-neonsme2.c)
 
-SET(NON_PROD_NEONSME2_MICROKERNEL_SRCS
-  src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c)
+SET(NON_PROD_NEONSME2_MICROKERNEL_SRCS)
 
 SET(ALL_NEONSME2_MICROKERNEL_SRCS ${PROD_NEONSME2_MICROKERNEL_SRCS} ${NON_PROD_NEONSME2_MICROKERNEL_SRCS})
diff --git a/gen/neonsme2_microkernels.bzl b/gen/neonsme2_microkernels.bzl
index 3024cc5b07..cecf32795c 100644
--- a/gen/neonsme2_microkernels.bzl
+++ b/gen/neonsme2_microkernels.bzl
@@ -15,6 +15,7 @@ PROD_NEONSME2_MICROKERNEL_SRCS = [
     "src/pqs8-f32-qc8w-igemm/pqs8-f32-qc8w-igemm-32x32c4-minmax-neonsme2.c",
     "src/pqs8-qc8w-gemm/pqs8-qc8w-gemm-1x32c4-minmax-neonsme2.c",
     "src/pqs8-qc8w-gemm/pqs8-qc8w-gemm-32x32c4-minmax-neonsme2.c",
+    "src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c",
     "src/qp8-f32-qc4w-gemm/qp8-f32-qc4w-gemm-minmax-1x64c4-neonsme2.c",
     "src/qp8-f32-qc4w-gemm/qp8-f32-qc4w-gemm-minmax-16x64c4-neonsme2.c",
     "src/qp8-f32-qc8w-gemm/qp8-f32-qc8w-gemm-minmax-1x64c4-neonsme2.c",
@@ -28,7 +29,6 @@ PROD_NEONSME2_MICROKERNEL_SRCS = [
 ]
 
 NON_PROD_NEONSME2_MICROKERNEL_SRCS = [
-    "src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c",
 ]
 
 ALL_NEONSME2_MICROKERNEL_SRCS = PROD_NEONSME2_MICROKERNEL_SRCS + NON_PROD_NEONSME2_MICROKERNEL_SRCS
diff --git a/src/configs/gemm-config.c b/src/configs/gemm-config.c
index 8c1f146a98..b7c9cd2fdf 100644
--- a/src/configs/gemm-config.c
+++ b/src/configs/gemm-config.c
@@ -2298,6 +2298,31 @@ static void init_qd8_f16_qc4w_gemm_config(void) {
     const struct xnn_hardware_config* hardware_config = xnn_init_hardware_config();
     assert(hardware_config != NULL);
     (void) hardware_config;  // May be unused.
+    #if XNN_ENABLE_KLEIDIAI && XNN_ENABLE_ARM_SME2
+    if (hardware_config->arch_flags & xnn_arch_arm_sme2) {
+      const size_t mr =
+          xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2_get_mr();
+      const size_t nr =
+          xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2_get_nr();
+      qd8_f16_qc4w_gemm_config.minmax.dqgemm[XNN_MR_TO_INDEX(1)] =
+          XNN_INIT_HMP_DQGEMM_UKERNEL(
+              xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2);
+      qd8_f16_qc4w_gemm_config.minmax.dqgemm[XNN_MR_TO_INDEX(mr)] =
+          XNN_INIT_HMP_DQGEMM_UKERNEL(
+              xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2);
+      qd8_f16_qc4w_gemm_config.init.f16_qc4w =
+          xnn_init_f16_qc4w_minmax_scalar_params;
+      qd8_f16_qc4w_gemm_config.pack_weights_and_biases =
+          xnn_pack_kai_qs4_weights_and_biases_sme;
+      qd8_f16_qc4w_gemm_config.packed_stride_weights_and_biases =
+          xnn_packed_stride_kai_qs4_weights_and_biases_sme;
+      qd8_f16_qc4w_gemm_config.mr = mr;
+      qd8_f16_qc4w_gemm_config.nr = nr;
+      qd8_f16_qc4w_gemm_config.log2_kr = 2;
+      qd8_f16_qc4w_gemm_config.log2_sr = 0;
+      qd8_f16_qc4w_gemm_config.planes = 2;
+    } else
+    #endif  // XNN_ENABLE_KLEIDIAI && XNN_ENABLE_ARM_SME2
     if (XNN_ENABLE_ARM_I8MM && (hardware_config->arch_flags & xnn_arch_arm_neon_i8mm)) {
       #if XNN_ENABLE_ARM_I8MM
         qd8_f16_qc4w_gemm_config.minmax.dqgemm[XNN_MR_TO_INDEX(1)] = XNN_INIT_HMP_DQGEMM_UKERNEL(xnn_qd8_f16_qc4w_gemm_minmax_ukernel_1x16c8__neoni8mm);
diff --git a/src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c b/src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c
index 951ff1c994..7625d5491e 100644
--- a/src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c
+++ b/src/qd8-f16-qc4w-gemm/qd8-f16-qc4w-gemm-minmax-16x64c4-neonsme2.c
@@ -5,6 +5,7 @@
 
 #include <stddef.h>
 #include <stdint.h>
+#include <stdlib.h>
 
 #include "src/xnnpack/math.h"
 #include "src/xnnpack/microparams.h"
@@ -61,3 +62,43 @@ static void pack_lhs(
     packed += mr * packed_row_size;
   }
 }
+
+void xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2(
+    size_t mr, size_t nr, size_t k, const int8_t* a, size_t a_stride,
+    const void* w, xnn_float16* c, size_t cm_stride, size_t cn_stride,
+    const struct xnn_f16_qc4w_minmax_params* params,
+    const struct xnn_qd8_quantization_params* quantization_params) {
+#if XNN_ENABLE_KLEIDIAI
+  const size_t packed_mr =
+      kai_get_mr_matmul_clamp_f16_qai8dxp1vlx8_qsi4cxp4vlx8_1vlx4vl_sme2_mopa();
+  const size_t kr =
+      kai_get_kr_matmul_clamp_f16_qai8dxp1vlx8_qsi4cxp4vlx8_1vlx4vl_sme2_mopa();
+  const size_t sr =
+      kai_get_sr_matmul_clamp_f16_qai8dxp1vlx8_qsi4cxp4vlx8_1vlx4vl_sme2_mopa();
+  assert(sr == 1);
+  assert(round_up(k, 32) % kr == 0);
+  const size_t packed_lhs_size =
+      lhs_packed_size(mr, k, packed_mr);
+  void* packed_lhs = malloc(packed_lhs_size);
+  assert(packed_lhs != NULL);
+  pack_lhs(mr, k, packed_mr, kr, a, a_stride, quantization_params, packed_lhs);
+  kai_run_matmul_clamp_f16_qai8dxp1vlx8_qsi4cxp4vlx8_1vlx4vl_sme2_mopa(
+      mr, nr, k, packed_lhs, w, c, cm_stride, cn_stride,
+      xnn_float16_to_float(params->scalar.min),
+      xnn_float16_to_float(params->scalar.max));
+  free(packed_lhs);
+#else
+  (void) mr;
+  (void) nr;
+  (void) k;
+  (void) a;
+  (void) a_stride;
+  (void) w;
+  (void) c;
+  (void) cm_stride;
+  (void) cn_stride;
+  (void) params;
+  (void) quantization_params;
+  assert(!"KleidiAI support is required for this microkernel");
+#endif  // XNN_ENABLE_KLEIDIAI
+}
diff --git a/src/xnnpack/gemm.h b/src/xnnpack/gemm.h
index 9fc4cfa5a8..1f5eddb5f4 100644
--- a/src/xnnpack/gemm.h
+++ b/src/xnnpack/gemm.h
@@ -4044,6 +4044,8 @@ size_t xnn_qp8_f32_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2_get_nr();
 
 size_t xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2_get_mr(void);
 size_t xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2_get_nr(void);
+DECLARE_QD8_F16_QC4W_GEMM_MINMAX_UKERNEL_FUNCTION(
+    xnn_qd8_f16_qc4w_gemm_minmax_ukernel_16x64c4__neonsme2)
 
 size_t xnn_qp8_f32_qc4w_gemm_minmax_ukernel_1x64c4__neonsme_get_mr();
 size_t xnn_qp8_f32_qc4w_gemm_minmax_ukernel_1x64c4__neonsme_get_nr();
-- 
2.43.0
