From d9078339d0bf1ef78990570a8fbc0e5b999aadb4 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sun, 27 Sep 2026 15:06:35 +0100
Subject: [PATCH 05/18] hip : use short vectors for bf16 WMMA builtins with
 LLVM 23

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-cuda/mma.cuh | 13 +++++++------
 1 file changed, 7 insertions(+), 6 deletions(-)

diff --git a/ggml/src/ggml-cuda/mma.cuh b/ggml/src/ggml-cuda/mma.cuh
index 3583ba5e1..fc35803aa 100644
--- a/ggml/src/ggml-cuda/mma.cuh
+++ b/ggml/src/ggml-cuda/mma.cuh
@@ -1318,19 +1318,20 @@ namespace ggml_cuda_mma {
     static __device__ __forceinline__ void mma(
             tile<16, 16, float, dl_d> & D, const tile<16, 8, nv_bfloat162, dl_ab> & A, const tile<16, 8, nv_bfloat162, dl_ab> & B) {
 #if defined(AMD_WMMA_AVAILABLE)
+        // The bf16 WMMA builtins take the bf16 bits as 16-bit integer vectors.
 #if defined(RDNA4)
-        using bf16x8_t = __attribute__((ext_vector_type(8))) __bf16;
+        using shortx8_t = __attribute__((ext_vector_type(8))) short;
         using floatx8_t = __attribute__((ext_vector_type(8))) float;
         floatx8_t& acc_frag = reinterpret_cast<floatx8_t&>(D.x[0]);
-        const bf16x8_t& a_frag = reinterpret_cast<const bf16x8_t&>(A.x[0]);
-        const bf16x8_t& b_frag = reinterpret_cast<const bf16x8_t&>(B.x[0]);
+        const shortx8_t& a_frag = reinterpret_cast<const shortx8_t&>(A.x[0]);
+        const shortx8_t& b_frag = reinterpret_cast<const shortx8_t&>(B.x[0]);
         acc_frag = __builtin_amdgcn_wmma_f32_16x16x16_bf16_w32_gfx12(a_frag, b_frag, acc_frag);
 #elif defined(RDNA3)
-        using bf16x16_t = __attribute__((ext_vector_type(16))) __bf16;
+        using shortx16_t = __attribute__((ext_vector_type(16))) short;
         using floatx8_t = __attribute__((ext_vector_type(8))) float;
         floatx8_t& acc_frag = reinterpret_cast<floatx8_t&>(D.x[0]);
-        const bf16x16_t& a_frag = reinterpret_cast<const bf16x16_t&>(A.x[0]);
-        const bf16x16_t& b_frag = reinterpret_cast<const bf16x16_t&>(B.x[0]);
+        const shortx16_t& a_frag = reinterpret_cast<const shortx16_t&>(A.x[0]);
+        const shortx16_t& b_frag = reinterpret_cast<const shortx16_t&>(B.x[0]);
         acc_frag = __builtin_amdgcn_wmma_f32_16x16x16_bf16_w32(a_frag, b_frag, acc_frag);
 #else
         GGML_UNUSED_VARS(D, A, B);
-- 
2.55.0

