From ec928150501c2572fec05cb949061672bb424914 Mon Sep 17 00:00:00 2001 From: shaofeiqi Date: Fri, 18 Sep 2026 10:50:15 -0700 Subject: [PATCH 1/9] opencl: add bin kernel `kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin` (#28678) * opencl: add A8 Q6_K non-MoE binary kernel * opencl: fix layout compatibility --- ggml/src/ggml-opencl/CMakeLists.txt | 1 + ggml/src/ggml-opencl/ggml-opencl.cpp | 275 +++++++++++++++++- .../gemv_noshuffle_q6_k_f32_32b_trans.cl | 128 ++++++++ 3 files changed, 388 insertions(+), 16 deletions(-) create mode 100644 ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_32b_trans.cl diff --git a/ggml/src/ggml-opencl/CMakeLists.txt b/ggml/src/ggml-opencl/CMakeLists.txt index 45a7075b2..53e938618 100644 --- a/ggml/src/ggml-opencl/CMakeLists.txt +++ b/ggml/src/ggml-opencl/CMakeLists.txt @@ -191,6 +191,7 @@ set(GGML_OPENCL_KERNELS gemv_noshuffle_q6_k_f32_tiled gemm_noshuffle_q6_k_f32 gemm_noshuffle_q6_k_f32_tiled + gemv_noshuffle_q6_k_f32_32b_trans gemv_noshuffle_q5_k_f32 gemm_noshuffle_q5_k_f32 mul diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 28cf6172c..1c26797b9 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -1246,6 +1246,8 @@ struct ggml_backend_opencl_context { cl_kernel kernel_gemv_noshuffle_q6_K_f32_mc3; // multi-column (N=3) verify GEMV cl_kernel kernel_gemm_noshuffle_q6_K_f32; cl_kernel kernel_gemm_noshuffle_q6_K_f32_cok; + cl_kernel kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin; + cl_kernel kernel_gemv_noshuffle_q6_k_f32_32b_trans; cl_kernel kernel_gemv_noshuffle_q5_k_f32; cl_kernel kernel_gemv_noshuffle_q5_k_f32_mc3; // multi-column (N=3) verify GEMV (spec/MTP) cl_kernel kernel_gemm_noshuffle_q5_k_f32; @@ -4367,6 +4369,43 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { } } + backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans = nullptr; + backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin = nullptr; + if (backend_ctx->adreno_gen == ADRENO_GPU_GEN::X2E) { + { + std::string opts = std::string("-cl-std=") + opencl_c_std + + " -cl-mad-enable " + " -DSIMDGROUP_WIDTH=" + + std::to_string(backend_ctx->adreno_wave_size); +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "gemv_noshuffle_q6_k_f32_32b_trans.cl.h" + }; +#else + const std::string kernel_src = read_file("gemv_noshuffle_q6_k_f32_32b_trans.cl"); +#endif + cl_program prog = build_program_from_source(backend_ctx, kernel_src.c_str(), opts); + CL_CHECK((backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans = + clCreateKernel(prog, "kernel_gemv_noshuffle_q6_k_f32_32b_trans", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + + if (use_adreno_bin_kernels(backend_ctx)) { + size_t bin_size = 0; + const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("gemm_noshuffle_q6_k_f32_32b_trans_ila_a8", &bin_size); + if (kernel_bin && bin_size > 0) { + cl_program bin_prog = + build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, "", bin_size); + + CL_CHECK((backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin = + clCreateKernel(bin_prog, "kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8", &err), err)); + CL_CHECK(clReleaseProgram(bin_prog)); + GGML_LOG_CONT("."); + } + } + } + std::string CL_moe_compile_opts = std::string("-cl-std=") + opencl_c_std + " -cl-mad-enable " " -cl-fast-relaxed-math"; @@ -7294,6 +7333,8 @@ struct ggml_tensor_extra_cl_q6_K { cl_mem ql_img = nullptr; // Upper 2 bits of quantized weights. cl_mem qh = nullptr; + // Upper 2 bits as image1d_buffer_t + cl_mem qh_img = nullptr; // Scales for each block. cl_mem s = nullptr; // Scales for each super block. @@ -7329,6 +7370,10 @@ struct ggml_tensor_extra_cl_q6_K { CL_CHECK(clReleaseMemObject(ql_img)); ql_img = nullptr; } + if (qh_img != nullptr) { + CL_CHECK(clReleaseMemObject(qh_img)); + qh_img = nullptr; + } size_ql = 0; size_qh = 0; @@ -8566,6 +8611,21 @@ static inline bool use_flat_gemv_for_large_m_q6_K(const ggml_backend_opencl_cont && tensor->ne[2] == 1 && tensor->ne[3] == 1; } +inline bool use_q6_k_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + if (!backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans || + !backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin) { + return false; + } + return (tensor->ne[0] % 256 == 0) && (tensor->ne[1] % 64 == 0) && + !use_q6k_tiled(backend_ctx, tensor) && !use_flat_gemv_for_large_m_q6_K(backend_ctx, tensor); +#else + GGML_UNUSED(backend_ctx); + GGML_UNUSED(tensor); + return false; +#endif +} + inline bool use_q4_k_bin_kernels(const ggml_backend_opencl_context *backend_ctx, const ggml_tensor *tensor) { #ifdef GGML_OPENCL_USE_ADRENO_KERNELS if (!backend_ctx->kernel_gemv_noshuffle_q4_k_f32_32b_trans || @@ -11181,18 +11241,39 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer, cl_int M = tensor->ne[1]; // ne01 cl_int K = tensor->ne[0]; // ne00 - // Transpose ql as ushort - transpose_2d_as_16b(backend_ctx, - extra->ql, extra->ql, size_ql, K/4, M); + if (use_q6_k_bin_kernels(backend_ctx, tensor)) { + GGML_ASSERT(K % 256 == 0); + GGML_ASSERT(M % 64 == 0); - // Transpose qh as uchar - transpose_2d_as_8b(backend_ctx, - extra->qh, extra->qh, size_qh, K/4, M); + transpose_2d_as_32b(backend_ctx, extra->ql, extra->ql, size_ql, K/8, M); + transpose_2d_as_32b(backend_ctx, extra->qh, extra->qh, size_qh, K/16, M); - // Transpose s as ushort - transpose_2d_as_16b(backend_ctx, - extra->s, extra->s, size_s, K/16/2, M); + cl_image_format wimg_fmt = { CL_R, CL_UNSIGNED_INT32 }; + cl_image_desc wimg_desc; + memset(&wimg_desc, 0, sizeof(wimg_desc)); + wimg_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + wimg_desc.image_width = static_cast(ggml_nelements(tensor) / 8); + wimg_desc.buffer = extra->ql; + CL_CHECK((extra->ql_img = clCreateImage(context, CL_MEM_READ_ONLY, &wimg_fmt, &wimg_desc, NULL, &err), err)); + memset(&wimg_desc, 0, sizeof(wimg_desc)); + wimg_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + wimg_desc.image_width = static_cast(ggml_nelements(tensor) / 16); + wimg_desc.buffer = extra->qh; + CL_CHECK((extra->qh_img = clCreateImage(context, CL_MEM_READ_ONLY, &wimg_fmt, &wimg_desc, NULL, &err), err)); + } else { + // Transpose ql as ushort + transpose_2d_as_16b(backend_ctx, + extra->ql, extra->ql, size_ql, K/4, M); + + // Transpose qh as uchar + transpose_2d_as_8b(backend_ctx, + extra->qh, extra->qh, size_qh, K/4, M); + + // Transpose s as ushort + transpose_2d_as_16b(backend_ctx, + extra->s, extra->s, size_s, K/16/2, M); + } // Transpose d as ushort transpose_2d_as_16b(backend_ctx, extra->d, extra->d, size_d, K/256, M); @@ -12317,15 +12398,24 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer, buf_trans_ql.allocate(backend_ctx->context, size_ql); buf_trans_qh.allocate(backend_ctx->context, size_qh); - buf_trans_s.allocate(backend_ctx->context, size_s); buf_trans_d.allocate(backend_ctx->context, size_d); buf_unpacked.allocate(backend_ctx->context, ggml_nbytes(tensor)); - // transpose ql, qh, s and d back - transpose_2d_as_16b(backend_ctx, extra->ql, buf_trans_ql.buffer, size_ql, M, K/4); - transpose_2d_as_8b(backend_ctx, extra->qh, buf_trans_qh.buffer, size_qh, M, K/4); - transpose_2d_as_16b(backend_ctx, extra->s, buf_trans_s.buffer, size_s, M, K/16/2); - transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256); + cl_mem s_buffer; + if (use_q6_k_bin_kernels(backend_ctx, tensor)) { + transpose_2d_as_32b(backend_ctx, extra->ql, buf_trans_ql.buffer, size_ql, M, K/8); + transpose_2d_as_32b(backend_ctx, extra->qh, buf_trans_qh.buffer, size_qh, M, K/16); + // s is left row-major, untransposed, for the binary layout. + s_buffer = extra->s; + } else { + // transpose ql, qh, s and d back + buf_trans_s.allocate(backend_ctx->context, size_s); + transpose_2d_as_16b(backend_ctx, extra->ql, buf_trans_ql.buffer, size_ql, M, K/4); + transpose_2d_as_8b(backend_ctx, extra->qh, buf_trans_qh.buffer, size_qh, M, K/4); + transpose_2d_as_16b(backend_ctx, extra->s, buf_trans_s.buffer, size_s, M, K/16/2); + s_buffer = buf_trans_s.buffer; + } + transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256); // unpack cl_uchar mask = 0xFF; @@ -12333,7 +12423,7 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer, cl_kernel kernel = backend_ctx->kernel_restore_block_q6_K_noshuffle; CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &buf_trans_ql.buffer)); CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &buf_trans_qh.buffer)); - CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &buf_trans_s.buffer)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &s_buffer)); CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &buf_trans_d.buffer)); CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &buf_unpacked.buffer)); CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_uchar), &mask)); @@ -21111,6 +21201,145 @@ static void ggml_cl_mul_mat_q4_k_f32_adreno(ggml_backend_t backend, const ggml_t #endif } +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS +static void ggml_cl_mul_mat_q6_K_f32_adreno_ila(ggml_backend_t backend, const ggml_tensor * src0, + const ggml_tensor * src1, ggml_tensor * dst) { + GGML_ASSERT(src0); + GGML_ASSERT(src0->extra); + GGML_ASSERT(src1); + GGML_ASSERT(src1->extra); + GGML_ASSERT(dst); + GGML_ASSERT(dst->extra); + + ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context; + + ggml_tensor_extra_cl_q6_K * extra0_q6_K = (ggml_tensor_extra_cl_q6_K *)src0->extra; + ggml_tensor_extra_cl * extra1 = (ggml_tensor_extra_cl *)src1->extra; + ggml_tensor_extra_cl * extrad = (ggml_tensor_extra_cl *)dst->extra; + + cl_ulong offset1 = extra1->offset + src1->view_offs; + cl_ulong offsetd = extrad->offset + dst->view_offs; + + const int ne00 = src0->ne[0]; + const int ne01 = src0->ne[1]; + + const int ne1 = dst->ne[1]; + + GGML_ASSERT(ne00 % ggml_blck_size(src0->type) == 0); + + cl_context context = backend_ctx->context; + cl_kernel kernel; + + cl_int err; + cl_buffer_region region; + cl_image_format img_fmt; + cl_image_desc img_desc; + + const int M = ne01; + const int N = ne1; + const int K = ne00; + + if (ne1 == 1) { + cl_mem b_sub_buf = nullptr; + cl_mem b_img = nullptr; + + region.origin = offset1; + region.size = (size_t)K * N * sizeof(float); + CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err)); + + img_fmt = { CL_RGBA, CL_FLOAT }; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = (size_t)K * N / 4; + img_desc.buffer = b_sub_buf; + CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + + kernel = backend_ctx->kernel_gemv_noshuffle_q6_k_f32_32b_trans; + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q6_K->ql_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q6_K->qh_img)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q6_K->s)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q6_K->d)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_int), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(cl_int), &ne01)); + + size_t local_work_size[3] = { 64, 8, 1 }; + size_t global_work_size[3] = { (size_t)ne01, 8, 1 }; + backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); + + CL_CHECK(clReleaseMemObject(b_img)); + CL_CHECK(clReleaseMemObject(b_sub_buf)); + } else { + const int gemm_tile_n = 64; + int N_pad = CEIL_DIV(N, gemm_tile_n) * gemm_tile_n; + + cl_mem b_sub_buf = nullptr; + cl_mem b_padded = nullptr; + cl_mem b_buf = nullptr; + if (N_pad == N) { + region.origin = offset1; + region.size = (size_t)K * N * sizeof(float); + CL_CHECK((b_sub_buf = clCreateSubBuffer(extra1->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err)); + b_buf = b_sub_buf; + } else { + CL_CHECK((b_padded = clCreateBuffer(context, CL_MEM_READ_WRITE, (size_t)K * N_pad * sizeof(float), NULL, &err), err)); + const float zero = 0.0f; + CL_CHECK(clEnqueueFillBuffer(backend_ctx->queue, b_padded, &zero, sizeof(zero), 0, (size_t)K * N_pad * sizeof(float), 0, NULL, NULL)); + CL_CHECK(clEnqueueCopyBuffer(backend_ctx->queue, extra1->data_device, b_padded, offset1, 0, (size_t)K * N * sizeof(float), 0, NULL, NULL)); + b_buf = b_padded; + } + + img_fmt = { CL_R, CL_FLOAT }; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = (size_t)K * N_pad; + img_desc.buffer = b_buf; + cl_mem b_img; + CL_CHECK((b_img = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + + region.origin = offsetd; + region.size = (size_t)M * N * sizeof(float); + cl_mem d_sub_buf; + CL_CHECK((d_sub_buf = clCreateSubBuffer(extrad->data_device, 0, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err), err)); + img_fmt = { CL_R, CL_FLOAT }; + memset(&img_desc, 0, sizeof(img_desc)); + img_desc.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc.image_width = (size_t)M * N; + img_desc.buffer = d_sub_buf; + cl_mem d_img; + CL_CHECK((d_img = clCreateImage(context, CL_MEM_WRITE_ONLY, &img_fmt, &img_desc, NULL, &err), err)); + + kernel = backend_ctx->kernel_gemm_noshuffle_q6_k_f32_32b_trans_ila_a8_bin; + CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &extra0_q6_K->ql_img)); + CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra0_q6_K->qh)); + CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &extra0_q6_K->s)); + CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &extra0_q6_K->d)); + CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &b_img)); + CL_CHECK(clSetKernelArg(kernel, 5, sizeof(cl_mem), &d_img)); + CL_CHECK(clSetKernelArg(kernel, 6, sizeof(cl_uint), &ne00)); + CL_CHECK(clSetKernelArg(kernel, 7, sizeof(cl_uint), &ne01)); + CL_CHECK(clSetKernelArg(kernel, 8, sizeof(int), &N)); + + size_t local_work_size[3] = { 64, 2, 2 }; + size_t m_tiles = (size_t)CEIL_DIV(M, 64); + size_t global_work_size[3] = { 64, m_tiles, (size_t)CEIL_DIV(N_pad, gemm_tile_n) }; + backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); + + CL_CHECK(clReleaseMemObject(b_img)); + if (b_sub_buf) { + CL_CHECK(clReleaseMemObject(b_sub_buf)); + } + if (b_padded) { + CL_CHECK(clReleaseMemObject(b_padded)); + } + CL_CHECK(clReleaseMemObject(d_img)); + CL_CHECK(clReleaseMemObject(d_sub_buf)); + } +} +#endif // GGML_OPENCL_USE_ADRENO_KERNELS + static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) { #ifdef GGML_OPENCL_USE_ADRENO_KERNELS GGML_ASSERT(src0); @@ -21159,6 +21388,20 @@ static void ggml_cl_mul_mat_q6_K_f32_adreno(ggml_backend_t backend, const ggml_t // (the #1 MTP bottleneck; mc3 above can't, it reads the noshuffle layout). const bool use_q6k_tiled_mc = q6k_mc3 && (ne1 == 3) && (ne01 >= 32768) && use_q6k_tiled(backend_ctx, src0); + const bool use_bin = use_q6_k_bin_kernels(backend_ctx, src0); + + if (use_bin) { + if (use_q6k_mc3 || use_q6k_tiled_mc) { + static bool warned = false; + if (!warned) { + GGML_LOG_WARN("ggml_opencl: GGML_OPENCL_Q6K_MC3 is bypassed by Q6_K binary kernels\n"); + warned = true; + } + } + ggml_cl_mul_mat_q6_K_f32_adreno_ila(backend, src0, src1, dst); + return; + } + if (ne1 == 1 || use_q6k_mc3 || use_q6k_tiled_mc) { cl_mem ql_img = nullptr; cl_mem qh_img = nullptr; diff --git a/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_32b_trans.cl b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_32b_trans.cl new file mode 100644 index 000000000..2e1e2d76d --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/gemv_noshuffle_q6_k_f32_32b_trans.cl @@ -0,0 +1,128 @@ +#pragma OPENCL EXTENSION cl_khr_fp16 : enable +#pragma OPENCL EXTENSION cl_khr_subgroups : enable +#pragma OPENCL EXTENSION cl_qcom_reqd_sub_group_size : enable + +#define QK_K 256 +#define N_SIMDGROUP 8 +#define SIMDGROUP_WIDTH 64 + +static inline float8 q6_k_to_fp32_packed8(ushort2 ql8, ushort qh8, float d_scale) { + float8 fp32x8; + fp32x8.s0 = ((float)(( ql8.s0 & 0x000F) | ((uint)((qh8 ) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s1 = ((float)((( ql8.s0 >> 4) & 0x000F) | ((uint)((qh8 >> 2) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s2 = ((float)((( ql8.s0 >> 8) & 0x000F) | ((uint)((qh8 >> 4) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s3 = ((float)((( ql8.s0 >> 12)& 0x000F) | ((uint)((qh8 >> 6) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s4 = ((float)(( ql8.s1 & 0x000F) | ((uint)((qh8 >> 8) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s5 = ((float)((( ql8.s1 >> 4) & 0x000F) | ((uint)((qh8 >>10) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s6 = ((float)((( ql8.s1 >> 8) & 0x000F) | ((uint)((qh8 >>12) & 0x3) << 4)) - 32.f) * d_scale; + fp32x8.s7 = ((float)((( ql8.s1 >> 12)& 0x000F) | ((uint)((qh8 >>14) & 0x3) << 4)) - 32.f) * d_scale; + return fp32x8; +} + +__attribute__((qcom_reqd_sub_group_size("half"))) +__kernel void kernel_gemv_noshuffle_q6_k_f32_32b_trans( + __read_only image1d_buffer_t src0_ql, + __read_only image1d_buffer_t src0_qh, + __global char * src0_s, + __global half * src0_d, + __read_only image1d_buffer_t src1, + __global float * dst, + ulong offsetd, + int ne00, + int ne01 +) { + uint i01 = get_global_id(0); + uint sgid = get_local_id(1); + uint slid = get_sub_group_local_id(); + + int num_superblocks = ne00 / QK_K; + int num_subblocks = ne00 / 32; // 2 sub-blocks of 16 processed per iter below + int scales_per_row = num_superblocks * 16; + + __private float sum = 0.0f; + + // Loop over 32-element groups (2 sub-blocks of 16 each), N_SIMDGROUP groups per iter. + for (uint ib = sgid; ib < num_subblocks; ib += N_SIMDGROUP) { + uint sb = ib / 8; // super-block index + uint j = ib % 8; // 32-element group within super-block (0..7) + + // Load d for this super-block. + half d_val = src0_d[sb * ne01 + i01]; + + // Load 2 sub-block scales (int8), one per 16 elements. + global const char * sc = src0_s + i01 * scales_per_row + sb * 16; + float scale0 = (float)d_val * (float)sc[j * 2]; + float scale1 = (float)d_val * (float)sc[j * 2 + 1]; + + // Load 4 uints of ql (32 elements, 4-bit each = 128 bits), column-major stride ne01. + uint ql_base = (ib * 4) * ne01 + i01; + uint4 regQL; + regQL.s0 = read_imageui(src0_ql, ql_base).x; + regQL.s1 = read_imageui(src0_ql, ql_base + ne01).x; + regQL.s2 = read_imageui(src0_ql, ql_base + ne01 * 2).x; + regQL.s3 = read_imageui(src0_ql, ql_base + ne01 * 3).x; + + // Load 2 uints of qh (32 elements, 2-bit each = 64 bits), column-major stride ne01. + uint qh_base = (ib * 2) * ne01 + i01; + uint2 regQH; + regQH.s0 = read_imageui(src0_qh, qh_base).x; + regQH.s1 = read_imageui(src0_qh, qh_base + ne01).x; + + // Load activations: 32 floats = 8 float4s. + uint y_offset = ib * 8; + + float4 y_local = (slid < 8) ? read_imagef(src1, (y_offset + slid)) : (float4)0.0f; + float4 y0 = sub_group_broadcast(y_local, 0); + float4 y1 = sub_group_broadcast(y_local, 1); + float4 y2 = sub_group_broadcast(y_local, 2); + float4 y3 = sub_group_broadcast(y_local, 3); + float4 y4v = sub_group_broadcast(y_local, 4); + float4 y5 = sub_group_broadcast(y_local, 5); + float4 y6 = sub_group_broadcast(y_local, 6); + float4 y7 = sub_group_broadcast(y_local, 7); + + // Dequantize elements 0..7 (scale0). + float8 fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s0), (ushort)(regQH.s0 & 0xFFFF), scale0); + + float4 acc = y0 * fp32x8.lo; + acc += y1 * fp32x8.hi; + + // Dequantize elements 8..15 (scale0). + fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s1), (ushort)(regQH.s0 >> 16), scale0); + + acc += y2 * fp32x8.lo; + acc += y3 * fp32x8.hi; + + // Dequantize elements 16..23 (scale1). + fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s2), (ushort)(regQH.s1 & 0xFFFF), scale1); + + acc += y4v * fp32x8.lo; + acc += y5 * fp32x8.hi; + + // Dequantize elements 24..31 (scale1). + fp32x8 = q6_k_to_fp32_packed8(as_ushort2(regQL.s3), (ushort)(regQH.s1 >> 16), scale1); + + acc += y6 * fp32x8.lo; + acc += y7 * fp32x8.hi; + + sum += ((acc.s0 + acc.s1) + (acc.s2 + acc.s3)); + } + + // reduction in local memory, assumes #subgroups=4 + __local float reduceLM[SIMDGROUP_WIDTH * (N_SIMDGROUP - 1)]; + if (sgid > 0) { + reduceLM[SIMDGROUP_WIDTH * (sgid - 1) + slid] = sum; + } + barrier(CLK_LOCAL_MEM_FENCE); + if (sgid == 0) { + for (uint i = 0; i < N_SIMDGROUP - 1; ++i) { + sum += reduceLM[SIMDGROUP_WIDTH * i + slid]; + } + } + + // 1 output per thread in subgroup 0 + if (sgid == 0) { + dst = dst + (offsetd >> 2); + dst[i01] = sum; + } +} From 18a04f09c24616898792bcfaa17f3550bdc78912 Mon Sep 17 00:00:00 2001 From: Todor Boinovski Date: Fri, 18 Sep 2026 13:15:08 -0700 Subject: [PATCH 2/9] hexagon: HMX flash-attention head_dim padding (support DK=DV=72) (#26539) Allow HMX flash-attention to run with head_dim not a multiple of 64 (e.g. SigLIP head_dim=72), by operating on DK/DV rounded up to 64 with zero-filled tail lanes. --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 13 +- ggml/src/ggml-hexagon/htp/flash-attn-ops.c | 101 +++++++++--- ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h | 176 +++++++++++++++++++-- tests/test-backend-ops.cpp | 4 + 4 files changed, 256 insertions(+), 38 deletions(-) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index 3f1495645..af8013b08 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -4000,7 +4000,9 @@ static bool ggml_hexagon_flash_attn_is_hmx_eligible( const uint32_t DK = q->ne[0]; const uint32_t DV = v->ne[0]; - if (DK % 64 != 0 || DV % 64 != 0) { + // Head dims that are not multiples of 64 are handled by internally padding to + // DK_pad/DV_pad = round_up(.,64) and zero-filling the tail lanes. + if (DK % 8 != 0 || DV % 8 != 0) { return false; } @@ -4073,8 +4075,13 @@ static bool ggml_hexagon_precompute_flash_attn_params( // Check HMX eligibility const struct ggml_tensor * sinks = op->src[4]; if (ggml_hexagon_flash_attn_is_hmx_eligible(sess, q, k, v, sinks)) { + // HMX tiles head_dim in units of 64; when DK/DV are not 64-aligned the kernel + // operates on padded dims with zero-filled tail lanes. VTCM budget and chunk-size + // are sized for the padded tiles. + const uint32_t DK_pad = hex_round_up(DK, 64); + const uint32_t DV_pad = hex_round_up(DV, 64); size_t Br = 0, Bc = 0; - int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK, DV, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0); + int ret = hmx_fa_find_chunk_size(&Br, &Bc, G, DK_pad, DV_pad, neq1, nek1, sess->vtcm_size, sess->n_threads, kparams->is_q_fp32 != 0); if (ret == 0) { kparams->kernel_type = HTP_FA_KERNEL_HMX; kparams->Br = Br; @@ -4084,7 +4091,7 @@ static bool ggml_hexagon_precompute_flash_attn_params( kparams->u.hmx.g_br = hex_align_up(G * Br, 32); kparams->u.hmx.pipeline = (kparams->n_kv_blocks >= 3 && sess->n_threads >= 2) ? 1 : 0; - kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK, DV, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0); + kparams->vtcm_size = hmx_fa_compute_vtcm_usage(G, DK_pad, DV_pad, Br, Bc, kparams->n_threads, kparams->u.hmx.pipeline != 0, kparams->is_q_fp32 != 0); const size_t row_vec_bytes = hex_align_up(Bc * sizeof(uint16_t), 256); kparams->u.hmx.row_buf_stride = row_vec_bytes / 128; // HVX vector is 128 bytes diff --git a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c index 8a1caba22..75422f420 100644 --- a/ggml/src/ggml-hexagon/htp/flash-attn-ops.c +++ b/ggml/src/ggml-hexagon/htp/flash-attn-ops.c @@ -108,6 +108,7 @@ struct hmx_fa_context { // Dimensions uint32_t DK, DV; + uint32_t DK_pad, DV_pad; // head_dim rounded up to 64 for HMX tiling uint32_t n_kv; // kv_len uint32_t n_kv_heads; // number of KV heads uint32_t n_heads; // number of Q heads @@ -652,7 +653,7 @@ static void fa_k_interleave_thread(unsigned int n, unsigned int i, void * data) hvx_dequantize_row_q8_0_f16(row_k, row_k, factx->DK); } } - hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles[args->buf_idx], (const __fp16 *) args->curr_k, total_rows, factx->DK, + hmx_interleave_rows_to_tiles(factx->vtcm_k_tiles[args->buf_idx], (const __fp16 *) args->curr_k, total_rows, factx->DK_pad, args->src_stride, start, end); htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_K_PREP, (uint16_t) (args->kv_start + start)); } @@ -706,7 +707,7 @@ static void fa_v_interleave_thread(unsigned int n, unsigned int i, void * data) hvx_dequantize_row_q8_0_f16(row_v, row_v, factx->DV); } } - hmx_interleave_cols_to_tiles(v_tiles_dst, (const __fp16 *) args->v_src, total_rows, factx->DV, + hmx_interleave_cols_to_tiles(v_tiles_dst, (const __fp16 *) args->v_src, total_rows, factx->DV_pad, args->src_stride, (uint32_t) args->n_col_tiles, start, end); htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_FA_V_PREP, (uint16_t) (args->kv_start + start)); } @@ -832,17 +833,22 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) { const uint32_t kv_head = args->kv_head; const uint32_t ib3 = args->ib3; - assert(factx->DK == factx->DV); - const bool use_q_dma = (factx->vtcm_q_dma != NULL); __fp16 * q_tiles = factx->vtcm_q_tiles; + const size_t DK_pad = factx->DK_pad; if (use_q_dma) { const size_t g_rows_end = hex_smin(end, n_rows_g); const uint32_t d_limit = factx->is_q_fp32 ? DK / 32 : DK / 64; uint8_t * q_flat = (uint8_t *) factx->vtcm_q_dma; - if (factx->is_q_fp32) { + if (DK_pad != DK) { + if (factx->is_q_fp32) { + hmx_fa_q_prep_fp32_pad(q_tiles, q_flat, start, end, g_rows_end, DK, DK_pad, G, args->n_rows_q, &factx->div_G, args->q_transposed); + } else { + hmx_fa_q_prep_fp16_pad(q_tiles, q_flat, start, end, g_rows_end, DK, DK_pad, G, args->n_rows_q, &factx->div_G, args->q_transposed); + } + } else if (factx->is_q_fp32) { switch (d_limit) { case 2: hmx_fa_q_prep_fp32_d2(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break; case 4: hmx_fa_q_prep_fp32_d4(q_tiles, q_flat, start, end, g_rows_end, DK, G, args->n_rows_q, &factx->div_G, args->q_transposed); break; @@ -858,7 +864,7 @@ static void fa_q_load_thread(unsigned int n, unsigned int i, void * data) { } else { // Fallback: direct-from-DDR/L2 path hmx_fa_q_prep_fallback(q_tiles, q->data, q->nb[1], q->nb[2], q->nb[3], - q_start, kv_head, ib3, start, end, n_rows_g, G, DK, factx->is_q_fp32, &factx->div_G); + q_start, kv_head, ib3, start, end, n_rows_g, G, DK, DK_pad, factx->is_q_fp32, &factx->div_G); } } @@ -952,6 +958,8 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) { const uint32_t kv_head = args->kv_head; const uint32_t ib3 = args->ib3; + const size_t DV_pad = factx->DV_pad; + size_t q_idx = fastdiv(start, &factx->div_G); size_t h_idx = fastmodulo(start, G, &factx->div_G); @@ -961,7 +969,7 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) { size_t r0 = r / HMX_FP16_TILE_N_ROWS; size_t r1 = r % HMX_FP16_TILE_N_ROWS; - const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV; + const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV_pad; for (uint32_t d = 0; d < DV / 32; ++d) { const HVX_Vector * in_tile = (const HVX_Vector *) (tile_row_base + d * HMX_FP16_TILE_N_ELMS); @@ -972,6 +980,16 @@ static void fa_o_store_thread_f32(unsigned int n, unsigned int i, void * data) { *(HVX_UVector *) (out + d * 32) = Q6_V_hi_W(vp); } } + // Ragged tail: DV not a multiple of 32 (e.g. 72 -> last 8 lanes). Partial vector-write + // for the remaining (DV % 32) floats. + const uint32_t d_tail = DV / 32; + const uint32_t rem = DV - d_tail * 32; + if (rem) { + const HVX_Vector * in_tile = (const HVX_Vector *) (tile_row_base + d_tail * HMX_FP16_TILE_N_ELMS); + HVX_VectorPair vp = hvx_vec_f16_to_f32_shuff(in_tile[r1 / 2]); + HVX_Vector vd = (r1 % 2 == 0) ? Q6_V_lo_W(vp) : Q6_V_hi_W(vp); + hvx_vec_store_u((void *) (out + d_tail * 32), rem * sizeof(float), vd); + } h_idx++; if (h_idx == G) { @@ -1006,6 +1024,9 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) { const uint32_t kv_head = args->kv_head; const uint32_t ib3 = args->ib3; + // O-tiles use the padded head dim (DV_pad); dst holds the real DV lanes. + const size_t DV_pad = factx->DV_pad; + size_t q_idx = fastdiv(start, &factx->div_G); size_t h_idx = fastmodulo(start, G, &factx->div_G); @@ -1015,7 +1036,7 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) { size_t r0 = r / HMX_FP16_TILE_N_ROWS; size_t r1 = r % HMX_FP16_TILE_N_ROWS; - const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV; + const __fp16 * tile_row_base = o_tile_src + r0 * HMX_FP16_TILE_N_ROWS * DV_pad; for (uint32_t d = 0; d < DV / 64; ++d) { const __fp16 * in_dtile = tile_row_base + d * HMX_FP16_TILE_N_ELMS * 2; @@ -1028,6 +1049,17 @@ static void fa_o_store_thread_f16(unsigned int n, unsigned int i, void * data) { *(HVX_UVector *) (out + d * 64) = Q6_V_hi_W(vp); } } + // Ragged tail when DV is not a multiple of 64. + const uint32_t d_tail = DV / 64; + const uint32_t rem = DV - d_tail * 64; + if (rem) { + const __fp16 * in_dtile = tile_row_base + d_tail * HMX_FP16_TILE_N_ELMS * 2; + const HVX_Vector * pv_in0 = ((const HVX_Vector *) in_dtile) + r1 / 2; + const HVX_Vector * pv_in1 = pv_in0 + 16; + HVX_VectorPair vp = Q6_W_vdeal_VVR(*pv_in1, *pv_in0, -2); + HVX_Vector vd = (r1 % 2 == 0) ? Q6_V_lo_W(vp) : Q6_V_hi_W(vp); + hvx_vec_store_u((void *) (out + d_tail * 64), rem * sizeof(__fp16), vd); + } h_idx++; if (h_idx == G) { @@ -1829,8 +1861,11 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { const uint32_t DK = neq0; const uint32_t DV = nev0; - // HMX requires head_dim to be multiple of 32 - if (DK % 32 != 0 || DV % 32 != 0) { + // HMX tiles head_dim in units of 64. head_dim need not be 64- (or 32-) aligned: + // we can operate on DK/DV rounded up to 64 with tail lanes [D, D_pad) zero-filled. + const uint32_t DK_pad = hex_round_up(DK, 64); + const uint32_t DV_pad = hex_round_up(DV, 64); + if (DK == 0 || DV == 0) { return HTP_STATUS_NO_SUPPORT; } @@ -1847,6 +1882,8 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { factx.n_threads = kparams->n_threads; factx.DK = DK; factx.DV = DV; + factx.DK_pad = DK_pad; + factx.DV_pad = DV_pad; factx.n_kv = nek1; factx.n_kv_heads = n_kv_heads; factx.n_heads = neq2; @@ -1905,16 +1942,18 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { // ======== VTCM allocation (GQA-aware) ======== // K/V row sizes drive the DMA descriptors (not the VTCM layout) and are used - // throughout the KV loop below. + // throughout the KV loop below. The DMA copies only the real DK/DV columns; the + // staging rows are padded to hold DK_pad/DV_pad columns (tail zero-filled below) + // so the HMX interleave/tile logic can operate on 64-aligned head dims. const size_t size_k_row = htp_tensor_get_row_size(k->type, DK); const size_t size_v_row = htp_tensor_get_row_size(v->type, DV); - const size_t size_k_row_padded = hex_round_up(DK * sizeof(__fp16), 128); - const size_t size_v_row_padded = hex_round_up(DV * sizeof(__fp16), 128); + const size_t size_k_row_padded = hex_round_up(DK_pad * sizeof(__fp16), 128); + const size_t size_v_row_padded = hex_round_up(DV_pad * sizeof(__fp16), 128); // Build the VTCM layout once (shared with the host estimator) and place every - // scratch buffer at its computed offset. + // scratch buffer at its computed offset. Padded head dims size the HMX tiles. struct hmx_fa_vtcm_layout L; - hmx_fa_vtcm_layout_build(&L, G, DK, DV, Br, Bc, n_threads, pipeline, factx.is_q_fp32); + hmx_fa_vtcm_layout_build(&L, G, DK_pad, DV_pad, Br, Bc, n_threads, pipeline, factx.is_q_fp32); if (L.total_bytes > ctx->vtcm_size) { return HTP_STATUS_VTCM_TOO_SMALL; @@ -1961,6 +2000,24 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { dma_cache_init(&factx.m_cache, (uint8_t *) factx.vtcm_mask_buf, L.m_buf_slot_bytes, HMX_FA_DMA_CACHE_SIZE); + // Head-dim padding: the K/V DMA staging buffers and the flat-Q buffer are laid out + // with padded row strides (size_{k,v,q}_row_padded, covering D_pad columns) but the + // DMA only writes the real D columns per row. Zero the whole staging buffers once up + // front so tail lanes [D, D_pad) stay zero for all KV blocks. No-op when already aligned. + if (DK_pad != DK || DV_pad != DV) { + const size_t k_buf_bytes = (size_t) factx.Bc * size_k_row_padded; + const size_t v_buf_bytes = (size_t) factx.Bc * size_v_row_padded; + hvx_splat_u8_a((char *) factx.vtcm_k_fp16[0], 0, k_buf_bytes); + hvx_splat_u8_a((char *) factx.vtcm_k_fp16[1], 0, k_buf_bytes); + hvx_splat_u8_a((char *) factx.vtcm_v_fp16[0], 0, v_buf_bytes); + hvx_splat_u8_a((char *) factx.vtcm_v_fp16[1], 0, v_buf_bytes); + // Flat-Q DMA scratch + if (factx.vtcm_q_dma) { + const size_t q_dma_bytes = hex_align_up(factx.g_br * DK * (factx.is_q_fp32 ? sizeof(float) : sizeof(__fp16)), 128); + hvx_splat_u8_a((char *) factx.vtcm_q_dma, 0, q_dma_bytes); + } + } + // ======== Initialize HMX output scales ======== hmx_init_column_scales(factx.vtcm_hmx_scales_id, Q6_V_vsplat_R(0x3c00)); // 1.0 hmx_init_column_scales(factx.vtcm_hmx_scales_qk, hvx_vec_splat_f16(factx.scale)); @@ -2072,7 +2129,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { qk_job[0].s_tiles = factx.vtcm_s_tiles[0]; qk_job[0].n_row_tiles = n_row_tiles; qk_job[0].n_col_tiles = hmx_ceil_div(kv_rows0, HMX_FP16_TILE_N_COLS); - qk_job[0].n_dot_tiles = DK / 32; + qk_job[0].n_dot_tiles = DK_pad / 32; qk_job[0].n_tiles_per_bc = n_tiles_per_bc; qk_job[0].hmx_scales = factx.vtcm_hmx_scales_qk; hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job[0])); @@ -2116,7 +2173,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { hmx_ceil_div(hex_smin(Bc, nek1 - (kv_blk - 1) * Bc), HMX_FP16_TILE_N_COLS); ou_job[prev_buf].n_row_tiles_g_br = n_row_tiles_g_br; ou_job[prev_buf].n_tiles_per_bc = n_tiles_per_bc; - ou_job[prev_buf].DV = DV; + ou_job[prev_buf].DV = DV_pad; hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job[prev_buf])); } @@ -2134,7 +2191,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { qk_job[next_buf].s_tiles = factx.vtcm_s_tiles[next_buf]; qk_job[next_buf].n_row_tiles = n_row_tiles; qk_job[next_buf].n_col_tiles = hmx_ceil_div(next_rows, HMX_FP16_TILE_N_COLS); - qk_job[next_buf].n_dot_tiles = DK / 32; + qk_job[next_buf].n_dot_tiles = DK_pad / 32; qk_job[next_buf].n_tiles_per_bc = n_tiles_per_bc; qk_job[next_buf].hmx_scales = factx.vtcm_hmx_scales_qk; hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_qk_dot_worker, &qk_job[next_buf])); @@ -2198,7 +2255,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { ou_job[0].n_col_tiles = last_cols; ou_job[0].n_row_tiles_g_br = n_row_tiles_g_br; ou_job[0].n_tiles_per_bc = n_tiles_per_bc; - ou_job[0].DV = DV; + ou_job[0].DV = DV_pad; hmx_queue_push(hmx_q, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job[0])); // Overlapped: run HVX build diag inv L while HMX is busy executing the update @@ -2246,7 +2303,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { qk_job.s_tiles = factx.vtcm_s_tiles[0]; qk_job.n_row_tiles = n_row_tiles; qk_job.n_col_tiles = n_col_tiles; - qk_job.n_dot_tiles = (size_t) (DK / 32); + qk_job.n_dot_tiles = (size_t) (DK_pad / 32); qk_job.n_tiles_per_bc = n_tiles_per_bc; qk_job.hmx_scales = factx.vtcm_hmx_scales_qk; @@ -2302,7 +2359,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { ou_job.n_col_tiles = n_col_tiles; ou_job.n_row_tiles_g_br = n_row_tiles_g_br; ou_job.n_tiles_per_bc = n_tiles_per_bc; - ou_job.DV = DV; + ou_job.DV = DV_pad; hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_update_worker, &ou_job)); if (kv_blk + 1 == factx.n_kv_blocks) { @@ -2380,7 +2437,7 @@ int hmx_flash_attn_ext(struct htp_ops_context * octx) { on_job.hmx_scales = factx.vtcm_hmx_scales_id; on_job.n_row_tiles = n_row_tiles; on_job.n_row_tiles_g_br = n_row_tiles_g_br; - on_job.DV = DV; + on_job.DV = DV_pad; hmx_queue_push(ctx->hmx_queue, hmx_queue_make_desc(hmx_fa_o_norm_worker, &on_job)); hmx_queue_pop(ctx->hmx_queue); } diff --git a/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h b/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h index d6795bf0b..8fd299795 100644 --- a/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h +++ b/ggml/src/ggml-hexagon/htp/hmx-fa-kernels.h @@ -495,12 +495,140 @@ static inline void hmx_fa_q_prep_fp16( } +// Head-dim-padded Q-prep (f32). Used when DK is not a multiple of 64. +static inline void hmx_fa_q_prep_fp32_pad(__fp16 * vtcm_q_tiles, + const uint8_t * temp_q_vtcm, + size_t start, + size_t end, + size_t g_rows_end, + size_t dk_in, + size_t dk_out, + size_t G, + size_t n_rows_q, + const struct fastdiv_values * div_G, + bool q_transposed) { + const uint32_t n_out_tiles = (uint32_t) (dk_out / 32); + for (size_t r = start; r < end; r += 2) { + size_t r0 = r / HMX_FP16_TILE_N_ROWS; + size_t r1 = r % HMX_FP16_TILE_N_ROWS; + __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * dk_out; + + if (r >= g_rows_end) { + for (uint32_t d = 0; d < n_out_tiles; ++d) { + ((HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS))[r1 / 2] = Q6_V_vzero(); + } + continue; + } + + const size_t q_idx0 = fastdiv(r + 0, div_G); + const size_t h_idx0 = fastmodulo(r + 0, G, div_G); + const size_t q_idx1 = fastdiv(r + 1, div_G); + const size_t h_idx1 = fastmodulo(r + 1, G, div_G); + + const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0); + const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1); + + const HVX_UVector * pv_in0 = (const HVX_UVector *) (temp_q_vtcm + offset0 * dk_in * sizeof(float)); + const HVX_UVector * pv_in1 = (r + 1 < g_rows_end) ? (const HVX_UVector *) (temp_q_vtcm + offset1 * dk_in * sizeof(float)) : NULL; + + for (uint32_t d = 0; d < n_out_tiles; ++d) { + HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS); + const size_t base_lane = (size_t) d * 32; + const size_t real_lanes = (base_lane < dk_in) ? hex_smin(32, dk_in - base_lane) : 0; + + if (real_lanes == 0) { + out_tile[r1 / 2] = Q6_V_vzero(); + continue; + } + + HVX_Vector v0 = pv_in0[d]; + HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero(); + if (real_lanes < 32) { + // Straddle tile: keep the first real_lanes floats, zero the padded tail so + // the packed f16 lanes beyond DK are zero. + const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(float))); + v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero()); + v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero()); + } + out_tile[r1 / 2] = hvx_vec_f32_to_f16_shuff(v0, v1); + } + } +} + +// Head-dim-padded Q-prep (f16). Used when DK is not a multiple of 64. +static inline void hmx_fa_q_prep_fp16_pad(__fp16 * vtcm_q_tiles, + const uint8_t * temp_q_vtcm, + size_t start, + size_t end, + size_t g_rows_end, + size_t dk_in, + size_t dk_out, + size_t G, + size_t n_rows_q, + const struct fastdiv_values * div_G, + bool q_transposed) { + const uint32_t n_out_pairs = (uint32_t) (dk_out / 64); + for (size_t r = start; r < end; r += 2) { + size_t r0 = r / HMX_FP16_TILE_N_ROWS; + size_t r1 = r % HMX_FP16_TILE_N_ROWS; + __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * dk_out; + + if (r >= g_rows_end) { + for (uint32_t d = 0; d < n_out_pairs; ++d) { + __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2; + HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2; + HVX_Vector * pv_out1 = pv_out0 + 16; + *pv_out0 = Q6_V_vzero(); + *pv_out1 = Q6_V_vzero(); + } + continue; + } + + const size_t q_idx0 = fastdiv(r + 0, div_G); + const size_t h_idx0 = fastmodulo(r + 0, G, div_G); + const size_t q_idx1 = fastdiv(r + 1, div_G); + const size_t h_idx1 = fastmodulo(r + 1, G, div_G); + + const size_t offset0 = q_transposed ? (h_idx0 * n_rows_q + q_idx0) : (q_idx0 * G + h_idx0); + const size_t offset1 = q_transposed ? (h_idx1 * n_rows_q + q_idx1) : (q_idx1 * G + h_idx1); + + const HVX_UVector * pv_in0 = (const HVX_UVector *) (temp_q_vtcm + offset0 * dk_in * sizeof(__fp16)); + const HVX_UVector * pv_in1 = (r + 1 < g_rows_end) ? (const HVX_UVector *) (temp_q_vtcm + offset1 * dk_in * sizeof(__fp16)) : NULL; + + for (uint32_t d = 0; d < n_out_pairs; ++d) { + __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2; + HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2; + HVX_Vector * pv_out1 = pv_out0 + 16; + + const size_t base_lane = (size_t) d * 64; + const size_t real_lanes = (base_lane < dk_in) ? hex_smin(64, dk_in - base_lane) : 0; + + if (real_lanes == 0) { + *pv_out0 = Q6_V_vzero(); + *pv_out1 = Q6_V_vzero(); + continue; + } + + HVX_Vector v0 = pv_in0[d]; + HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero(); + if (real_lanes < 64) { + const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(__fp16))); + v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero()); + v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero()); + } + HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2); + *pv_out0 = Q6_V_lo_W(vp); + *pv_out1 = Q6_V_hi_W(vp); + } + } +} + static inline void hmx_fa_q_prep_fallback( __fp16 * vtcm_q_tiles, uintptr_t q_data, size_t q_nb1, size_t q_nb2, size_t q_nb3, uint32_t q_start, uint32_t kv_head, uint32_t ib3, size_t start, size_t end, size_t n_rows_g, - size_t G, size_t DK, bool is_q_fp32, + size_t G, size_t dk_in, size_t dk_out, bool is_q_fp32, const struct fastdiv_values * div_G ) { for (size_t r = start; r < end; r += 2) { @@ -518,33 +646,55 @@ static inline void hmx_fa_q_prep_fallback( size_t r0 = r / HMX_FP16_TILE_N_ROWS; size_t r1 = r % HMX_FP16_TILE_N_ROWS; - __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * DK; + __fp16 * out_base = vtcm_q_tiles + r0 * HMX_FP16_TILE_N_ROWS * dk_out; if (is_q_fp32) { const HVX_UVector * pv_in0 = q_ptr0 ? (const HVX_UVector *) q_ptr0 : NULL; const HVX_UVector * pv_in1 = q_ptr1 ? (const HVX_UVector *) q_ptr1 : NULL; - for (uint32_t d = 0; d < DK / 32; ++d) { - HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero(); - HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero(); - HVX_Vector v_hf = hvx_vec_f32_to_f16_shuff(v0, v1); + for (uint32_t d = 0; d < dk_out / 32; ++d) { + HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS); + const size_t base_lane = (size_t) d * 32; + const size_t real_lanes = (base_lane < dk_in) ? hex_smin(32, dk_in - base_lane) : 0; - HVX_Vector * out_tile = (HVX_Vector *) (out_base + d * HMX_FP16_TILE_N_ELMS); - out_tile[r1 / 2] = v_hf; + if (real_lanes == 0) { + out_tile[r1 / 2] = Q6_V_vzero(); + continue; + } + HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero(); + HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero(); + if (real_lanes < 32) { + const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(float))); + v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero()); + v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero()); + } + out_tile[r1 / 2] = hvx_vec_f32_to_f16_shuff(v0, v1); } } else { const HVX_UVector * pv_in0 = q_ptr0 ? (const HVX_UVector *) q_ptr0 : NULL; const HVX_UVector * pv_in1 = q_ptr1 ? (const HVX_UVector *) q_ptr1 : NULL; - for (uint32_t d = 0; d < DK / 64; ++d) { - HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero(); - HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero(); - HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2); - + for (uint32_t d = 0; d < dk_out / 64; ++d) { __fp16 * out_dtile = out_base + d * HMX_FP16_TILE_N_ELMS * 2; HVX_Vector * pv_out0 = ((HVX_Vector *) out_dtile) + r1 / 2; HVX_Vector * pv_out1 = pv_out0 + 16; + const size_t base_lane = (size_t) d * 64; + const size_t real_lanes = (base_lane < dk_in) ? hex_smin(64, dk_in - base_lane) : 0; + + if (real_lanes == 0) { + *pv_out0 = Q6_V_vzero(); + *pv_out1 = Q6_V_vzero(); + continue; + } + HVX_Vector v0 = pv_in0 ? pv_in0[d] : Q6_V_vzero(); + HVX_Vector v1 = pv_in1 ? pv_in1[d] : Q6_V_vzero(); + if (real_lanes < 64) { + const HVX_VectorPred keep = Q6_Q_vsetq_R((uint32_t) (real_lanes * sizeof(__fp16))); + v0 = Q6_V_vmux_QVV(keep, v0, Q6_V_vzero()); + v1 = Q6_V_vmux_QVV(keep, v1, Q6_V_vzero()); + } + HVX_VectorPair vp = Q6_W_vshuff_VVR(v1, v0, -2); *pv_out0 = Q6_V_lo_W(vp); *pv_out1 = Q6_V_hi_W(vp); } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 30792e409..d8f4c3708 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -10713,6 +10713,10 @@ static std::vector> make_test_cases_eval() { } } + // asymmetric head_dim (hsk != hsv) with one or both sides not 64-aligned + test_cases.emplace_back(new test_flash_attn_ext(72, 64, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + test_cases.emplace_back(new test_flash_attn_ext(64, 72, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16)); + // mixed quant and Q1_0 test cases test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q4_0)); test_cases.emplace_back(new test_flash_attn_ext(64, 64, 4, {1, 1}, 128, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_F16)); From 50631b3d2c569ad8e5c112090cd28570b1268ee0 Mon Sep 17 00:00:00 2001 From: Todor Boinovski Date: Fri, 18 Sep 2026 14:20:48 -0700 Subject: [PATCH 3/9] hexagon: im2col update (#29103) * ggml-hexagon: accept 1D and padded IM2COL ops * ggml-hexagon: make pure-DDR IM2COL kernel is_2D-aware * ggml-hexagon: extend IM2COL DMA patch-embed fast path to 1D * ggml-hexagon: add blocked-staging general IM2COL DMA kernel --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 12 - ggml/src/ggml-hexagon/htp/im2col-ops.c | 491 +++++++++++++++++++------ 2 files changed, 369 insertions(+), 134 deletions(-) diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index af8013b08..f6f2fdd28 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -5530,11 +5530,6 @@ static bool ggml_hexagon_supported_im2col(const struct ggml_hexagon_session * se const struct ggml_tensor * src1 = op->src[1]; const struct ggml_tensor * dst = op; - const bool is_2D = ((const int32_t *) op->op_params)[6] == 1; - if (!is_2D) { - return false; - } - // For now support F32->F32 and F32->F16 only. if (src1->type != GGML_TYPE_F32 || (dst->type != GGML_TYPE_F16 && dst->type != GGML_TYPE_F32)) { return false; @@ -5544,13 +5539,6 @@ static bool ggml_hexagon_supported_im2col(const struct ggml_hexagon_session * se return false; } - // For now keep padded OPs on CPU. Will revisit once we expand coverage past patch-embed OPs. - const int32_t p0 = ((const int32_t *) op->op_params)[2]; - const int32_t p1 = ((const int32_t *) op->op_params)[3]; - if (p0 != 0 || p1 != 0) { - return false; - } - GGML_UNUSED(sess); return true; } diff --git a/ggml/src/ggml-hexagon/htp/im2col-ops.c b/ggml/src/ggml-hexagon/htp/im2col-ops.c index 52bbc37d1..26af14ed5 100644 --- a/ggml/src/ggml-hexagon/htp/im2col-ops.c +++ b/ggml/src/ggml-hexagon/htp/im2col-ops.c @@ -25,17 +25,20 @@ struct htp_im2col_context { uint32_t npatches; // number of patches assigned to this dev uint32_t npatches_per_thread; // patches = N*OH*OW (pure-DDR kernel) - uint32_t pe_row_base; // first N*OH row index assigned to this dev (DMA path) - uint32_t pe_nrows; // number of N*OH rows assigned to this dev (DMA path) - uint32_t pe_rows_per_thread; // N*OH rows per worker - uint32_t pe_src_row_bytes; // one output row's source: IC*KH*IW*4, rounded 256 - uint32_t pe_dst_row_bytes; // one output row's dst: OW*patch_stride*2, rounded 256 + uint32_t pe_row_base; // first N*OH row index assigned to this dev (DMA path) + uint32_t pe_nrows; // number of N*OH rows assigned to this dev (DMA path) + uint32_t pe_rows_per_thread; // N*OH rows per worker + uint32_t pe_src_row_bytes; // one output row's source: IC*KH*IW*4, rounded 256 + uint32_t pe_dst_row_bytes; // one output row's dst: OW*patch_stride*2, rounded 256 // Patch-embed DMA path VTCM ping-pong. uint8_t * pe_vtcm_src; // base of the 2x src buffers region uint8_t * pe_vtcm_dst; // base of the 2x dst buffers region uint32_t pe_src_size_per_thread; // 2 * pe_src_row_bytes uint32_t pe_dst_size_per_thread; // 2 * pe_dst_row_bytes + + uint32_t pe_owb; // output-col block size + uint32_t pe_wb; // staged source window width }; // Per-op VTCM layout for the patch-embed DMA path @@ -59,83 +62,253 @@ static inline void htp_im2col_vtcm_layout_build(struct htp_im2col_vtcm_layout * L->total_bytes = L->off_dst + L->dst_bytes_per_thread * n_threads; } -#define IM2COL_PATCHEMBED_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \ - static void FNAME(unsigned int nth, unsigned int ith, void * data) { \ - struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \ - struct htp_ops_context * octx = ictx->octx; \ - struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \ - const struct htp_tensor * restrict src0 = octx->src[0]; \ - const struct htp_tensor * restrict src1 = octx->src[1]; \ - const struct htp_tensor * restrict dst = octx->dst; \ - const int32_t s0 = octx->op_params[0], s1 = octx->op_params[1]; \ - const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3]; \ - const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5]; \ - const uint32_t N = src1->ne[3], IC = src1->ne[2], IH = src1->ne[1], IW = src1->ne[0]; \ - const uint32_t KH = src0->ne[1], KW = src0->ne[0]; \ - const uint32_t OH = dst->ne[2]; \ - const uint32_t OW = dst->ne[1]; \ - const uint32_t patch_stride = IC * KH * KW; \ - const float * restrict src_data = (const float *) src1->data; \ - DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \ - const uint32_t patch_end = ictx->patch_base + ictx->npatches; \ - const uint32_t patch_start = ictx->patch_base + ictx->npatches_per_thread * ith; \ - const uint32_t patch_stop = MIN(patch_start + ictx->npatches_per_thread, patch_end);\ - if (patch_start >= patch_stop) { \ - return; \ - } \ - htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \ - for (uint32_t p = patch_start; p < patch_stop; p++) { \ - const uint32_t iow = p % OW; \ - const uint32_t ioh = (p / OW) % OH; \ - const uint32_t in = p / (OW * OH); \ - DST_CTYPE * restrict dst_patch = dst_data + (uint64_t) p * patch_stride; \ - for (uint32_t iic = 0; iic < IC; iic++) { \ - const float * restrict src_plane = src_data + ((uint64_t) in * IC + iic) * IH * IW; \ - for (uint32_t ikh = 0; ikh < KH; ikh++) { \ - const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \ - DST_CTYPE * restrict out_run = dst_patch + iic * (KH * KW) + ikh * KW; \ - if (iih < 0 || iih >= (int32_t) IH) { \ - SPLAT_FN(out_run, 0.0f, KW); \ - continue; \ - } \ - const int32_t iiw0 = (int32_t) iow * s0 - p0; \ - const float * restrict src_run = src_plane + (uint64_t) iih * IW + iiw0; \ - if (d0 == 1) { \ - /* contiguous source run: [lo,hi) is in-bounds, tails are zero pad */ \ - const int32_t lo = iiw0 < 0 ? -iiw0 : 0; \ - int32_t hi = (int32_t) IW - iiw0; \ - if (hi > (int32_t) KW) { \ - hi = (int32_t) KW; \ - } \ - if (hi <= lo) { \ - SPLAT_FN(out_run, 0.0f, KW); \ - } else { \ - if (lo > 0) { \ - SPLAT_FN(out_run, 0.0f, (uint32_t) lo); \ - } \ - COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (src_run + lo), \ - (uint32_t) (hi - lo)); \ - if (hi < (int32_t) KW) { \ - SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi)); \ - } \ - } \ - continue; \ - } \ - for (uint32_t ikw = 0; ikw < KW; ikw++) { \ - const int32_t iiw = (int32_t) iow * s0 + (int32_t) ikw * d0 - p0; \ - out_run[ikw] = (iiw < 0 || iiw >= (int32_t) IW) ? \ - (DST_CTYPE) 0.0f : \ - (DST_CTYPE) src_plane[(uint64_t) iih * IW + iiw]; \ - } \ - } \ - } \ - } \ - htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \ +#define IM2COL_PATCHEMBED_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \ + static void FNAME(unsigned int nth, unsigned int ith, void * data) { \ + struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \ + struct htp_ops_context * octx = ictx->octx; \ + struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \ + const struct htp_tensor * restrict src0 = octx->src[0]; \ + const struct htp_tensor * restrict src1 = octx->src[1]; \ + const struct htp_tensor * restrict dst = octx->dst; \ + const int32_t s0 = octx->op_params[0]; \ + const int32_t s1 = octx->op_params[1]; \ + const int32_t p0 = octx->op_params[2]; \ + const int32_t p1 = octx->op_params[3]; \ + const int32_t d0 = octx->op_params[4]; \ + const int32_t d1 = octx->op_params[5]; \ + const int32_t is_2D = octx->op_params[6] == 1; \ + const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; \ + const uint32_t IC = is_2D ? src1->ne[2] : src1->ne[1]; \ + const uint32_t IH = is_2D ? src1->ne[1] : 1; \ + const uint32_t IW = src1->ne[0]; \ + const uint32_t KH = is_2D ? src0->ne[1] : 1; \ + const uint32_t KW = src0->ne[0]; \ + const uint32_t OH = is_2D ? dst->ne[2] : 1; \ + const uint32_t OW = dst->ne[1]; \ + const uint32_t patch_stride = IC * KH * KW; \ + const float * restrict src_data = (const float *) src1->data; \ + DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \ + const uint32_t patch_end = ictx->patch_base + ictx->npatches; \ + const uint32_t patch_start = ictx->patch_base + ictx->npatches_per_thread * ith; \ + const uint32_t patch_stop = MIN(patch_start + ictx->npatches_per_thread, patch_end); \ + if (patch_start >= patch_stop) { \ + return; \ + } \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \ + for (uint32_t p = patch_start; p < patch_stop; p++) { \ + const uint32_t iow = p % OW; \ + const uint32_t ioh = (p / OW) % OH; \ + const uint32_t in = p / (OW * OH); \ + DST_CTYPE * restrict dst_patch = dst_data + (uint64_t) p * patch_stride; \ + for (uint32_t iic = 0; iic < IC; iic++) { \ + const float * restrict src_plane = src_data + ((uint64_t) in * IC + iic) * IH * IW; \ + for (uint32_t ikh = 0; ikh < KH; ikh++) { \ + const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \ + DST_CTYPE * restrict out_run = dst_patch + iic * (KH * KW) + ikh * KW; \ + if (iih < 0 || iih >= (int32_t) IH) { \ + SPLAT_FN(out_run, 0.0f, KW); \ + continue; \ + } \ + const int32_t iiw0 = (int32_t) iow * s0 - p0; \ + const float * restrict src_run = src_plane + (uint64_t) iih * IW + iiw0; \ + if (d0 == 1) { \ + /* contiguous source run: [lo,hi) is in-bounds, tails are zero pad */ \ + const int32_t lo = iiw0 < 0 ? -iiw0 : 0; \ + int32_t hi = (int32_t) IW - iiw0; \ + if (hi > (int32_t) KW) { \ + hi = (int32_t) KW; \ + } \ + if (hi <= lo) { \ + SPLAT_FN(out_run, 0.0f, KW); \ + } else { \ + if (lo > 0) { \ + SPLAT_FN(out_run, 0.0f, (uint32_t) lo); \ + } \ + COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (src_run + lo), \ + (uint32_t) (hi - lo)); \ + if (hi < (int32_t) KW) { \ + SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi)); \ + } \ + } \ + continue; \ + } \ + for (uint32_t ikw = 0; ikw < KW; ikw++) { \ + const int32_t iiw = (int32_t) iow * s0 + (int32_t) ikw * d0 - p0; \ + out_run[ikw] = (iiw < 0 || iiw >= (int32_t) IW) ? \ + (DST_CTYPE) 0.0f : \ + (DST_CTYPE) src_plane[(uint64_t) iih * IW + iiw]; \ + } \ + } \ + } \ + } \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, patch_start); \ } IM2COL_PATCHEMBED_BODY(im2col_patchembed_thread, __fp16, hvx_copy_f16_f32_uu, hvx_splat_f16_u, sizeof(__fp16), "f32-f16") IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx_splat_f32_u, sizeof(float), "f32-f32") +// Software-pipelined 2-deep: while HVX computes block bi from buffer slot +// (bi&1), the DMA engine stages block bi+1 into the other slot concurrently. +// A single dma_queue_flush per iteration (after issuing the next stage-in and +// this block's store-out) waits for both - safe because the ring is strict +// FIFO and each buffer slot is only reused after its prior consumer (compute +// or store-out) already finished in program order. +#define IM2COL_BLOCKED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \ + static void FNAME(unsigned int nth, unsigned int ith, void * data) { \ + struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \ + struct htp_ops_context * octx = ictx->octx; \ + struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \ + const struct htp_tensor * restrict src1 = octx->src[1]; \ + const struct htp_tensor * restrict dst = octx->dst; \ + const int32_t s0 = octx->op_params[0], s1 = octx->op_params[1]; \ + const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3]; \ + const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5]; \ + const int32_t is_2D = octx->op_params[6] == 1; \ + const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; \ + const uint32_t IC = is_2D ? src1->ne[2] : src1->ne[1]; \ + const uint32_t IH = is_2D ? src1->ne[1] : 1; \ + const uint32_t IW = src1->ne[0]; \ + const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; \ + const uint32_t KW = octx->src[0]->ne[0]; \ + const uint32_t OH = is_2D ? dst->ne[2] : 1; \ + const uint32_t OW = dst->ne[1]; \ + const uint32_t owb = ictx->pe_owb, Wb = ictx->pe_wb; \ + const uint32_t patch_stride = IC * KH * KW; \ + const float * restrict src_data = (const float *) src1->data; \ + DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \ + dma_queue * dmaq = octx->ctx->dma[ith]; \ + uint8_t * srcb_base = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread; \ + uint8_t * dstb_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \ + float * srcb2[2] = { (float *) srcb_base, (float *) (srcb_base + ictx->pe_src_row_bytes) }; \ + DST_CTYPE * dstb2[2] = { (DST_CTYPE *) dstb_base, (DST_CTYPE *) (dstb_base + ictx->pe_dst_row_bytes) }; \ + const uint32_t nrows = N * OH; \ + const uint32_t per_thread = ictx->pe_rows_per_thread; \ + const uint32_t row_start = per_thread * ith; \ + const uint32_t row_end = MIN(row_start + per_thread, nrows); \ + if (row_start >= row_end) \ + return; \ + const uint32_t nbpr = (OW + owb - 1) / owb; \ + const uint32_t nrows_local = row_end - row_start; \ + const uint32_t total_blocks = nrows_local * nbpr; \ + for (uint32_t bi = 0; bi < total_blocks; bi++) { \ + const uint32_t buf = bi & 1u; \ + float * srcb = srcb2[buf]; \ + DST_CTYPE * dstb = dstb2[buf]; \ + const uint32_t r = row_start + bi / nbpr; \ + const uint32_t in = r / OH; \ + const uint32_t ioh = r % OH; \ + const uint32_t c0 = (bi % nbpr) * owb; \ + const uint32_t nb = MIN(owb, OW - c0); \ + const int32_t win0 = (int32_t) c0 * s0 - p0; \ + if (bi == 0) { \ + /* prologue: stage block 0 and wait - nothing to overlap with yet */ \ + for (uint32_t ikh = 0; ikh < KH; ikh++) { \ + const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \ + if (iih < 0 || iih >= (int32_t) IH) \ + continue; \ + const int32_t lo = win0 < 0 ? -win0 : 0; \ + int32_t hi = (int32_t) IW - win0; \ + if (hi > (int32_t) Wb) \ + hi = (int32_t) Wb; \ + if (hi <= lo) \ + continue; \ + const uint32_t cpw = (uint32_t) (hi - lo); \ + float * vdst = srcb + (uint64_t) ikh * Wb + (uint32_t) lo; \ + const float * vsrc = src_data + ((uint64_t) (in * IC) * IH + iih) * IW + (win0 + lo); \ + while (!dma_queue_push(dmaq, dma_make_ptr((uint8_t *) vdst, (const uint8_t *) vsrc), \ + (size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float), \ + cpw * sizeof(float), IC)) { \ + dma_queue_pop(dmaq); \ + } \ + } \ + dma_queue_flush(dmaq); \ + } \ + if (bi + 1 < total_blocks) { \ + /* prefetch: stage block bi+1 into the other slot; overlaps with this block's compute below */ \ + const uint32_t nbuf = 1u - buf; \ + float * nsrcb = srcb2[nbuf]; \ + const uint32_t nr = row_start + (bi + 1) / nbpr; \ + const uint32_t nin = nr / OH; \ + const uint32_t nioh = nr % OH; \ + const uint32_t nc0 = ((bi + 1) % nbpr) * owb; \ + const int32_t nwin0 = (int32_t) nc0 * s0 - p0; \ + for (uint32_t ikh = 0; ikh < KH; ikh++) { \ + const int32_t iih = (int32_t) nioh * s1 + (int32_t) ikh * d1 - p1; \ + if (iih < 0 || iih >= (int32_t) IH) \ + continue; \ + const int32_t lo = nwin0 < 0 ? -nwin0 : 0; \ + int32_t hi = (int32_t) IW - nwin0; \ + if (hi > (int32_t) Wb) \ + hi = (int32_t) Wb; \ + if (hi <= lo) \ + continue; \ + const uint32_t cpw = (uint32_t) (hi - lo); \ + float * vdst = nsrcb + (uint64_t) ikh * Wb + (uint32_t) lo; \ + const float * vsrc = src_data + ((uint64_t) (nin * IC) * IH + iih) * IW + (nwin0 + lo); \ + while (!dma_queue_push(dmaq, dma_make_ptr((uint8_t *) vdst, (const uint8_t *) vsrc), \ + (size_t) KH * Wb * sizeof(float), (size_t) IH * IW * sizeof(float), \ + cpw * sizeof(float), IC)) { \ + dma_queue_pop(dmaq); \ + } \ + } \ + } \ + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, r); \ + for (uint32_t j = 0; j < nb; j++) { \ + const uint32_t iow = c0 + j; \ + DST_CTYPE * dst_patch = dstb + (uint64_t) j * patch_stride; \ + const int32_t iiw0 = (int32_t) iow * s0 - p0; \ + for (uint32_t ikh = 0; ikh < KH; ikh++) { \ + const int32_t iih = (int32_t) ioh * s1 + (int32_t) ikh * d1 - p1; \ + const int okh = (iih >= 0 && iih < (int32_t) IH); \ + for (uint32_t iic = 0; iic < IC; iic++) { \ + DST_CTYPE * out_run = dst_patch + iic * (KH * KW) + ikh * KW; \ + if (!okh) { \ + SPLAT_FN(out_run, 0.0f, KW); \ + continue; \ + } \ + const float * vrow = srcb + ((uint64_t) (iic * KH + ikh)) * Wb; /* col win0 at idx 0*/ \ + if (d0 == 1) { \ + /* contiguous run within the staged window: [lo,hi) in-bounds, tails zero pad */ \ + const int32_t lo = iiw0 < 0 ? -iiw0 : 0; \ + int32_t hi = (int32_t) IW - iiw0; \ + if (hi > (int32_t) KW) { \ + hi = (int32_t) KW; \ + } \ + if (hi <= lo) { \ + SPLAT_FN(out_run, 0.0f, KW); \ + } else { \ + if (lo > 0) { \ + SPLAT_FN(out_run, 0.0f, (uint32_t) lo); \ + } \ + COPY_FN((uint8_t *) (out_run + lo), (const uint8_t *) (vrow + (iiw0 + lo - win0)), \ + (uint32_t) (hi - lo)); \ + if (hi < (int32_t) KW) { \ + SPLAT_FN(out_run + hi, 0.0f, (KW - (uint32_t) hi)); \ + } \ + } \ + continue; \ + } \ + for (uint32_t ikw = 0; ikw < KW; ikw++) { \ + const int32_t iiw = iiw0 + (int32_t) ikw * d0; \ + out_run[ikw] = \ + (iiw < 0 || iiw >= (int32_t) IW) ? (DST_CTYPE) 0.0f : (DST_CTYPE) vrow[iiw - win0]; \ + } \ + } \ + } \ + } \ + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, r); \ + DST_CTYPE * ddr = dst_data + ((uint64_t) (in * OH + ioh) * OW + c0) * patch_stride; \ + dma_queue_push_vtcm_to_ddr(dmaq, dma_make_ptr((uint8_t *) ddr, (uint8_t *) dstb), \ + nb * patch_stride * (DST_ELEM), nb * patch_stride * (DST_ELEM), 1); \ + dma_queue_flush(dmaq); \ + } \ + } +IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_thread, __fp16, hvx_copy_f16_f32_uu, hvx_splat_f16_u, sizeof(__fp16), "blk-dma-f16") +IM2COL_BLOCKED_DMA_BODY(im2col_blocked_dma_f32_thread, float, hvx_copy_f32_uu, hvx_splat_f32_u, sizeof(float), "blk-dma-f32") + +// Exact-tiling patch-embed DMA fast path (s0==KW, p0=0, d0=1; and 2D s1==KH, +// p1=0, d1=1). Intentionally reads no stride/pad/dilation params so the inner +// copy stays tight and fully hoisted - do NOT graft the general gather in here. #define IM2COL_PATCHEMBED_DMA_BODY(FNAME, DST_CTYPE, COPY_FN, SPLAT_FN, DST_ELEM, TAG) \ static void FNAME(unsigned int nth, unsigned int ith, void * data) { \ struct htp_im2col_context * ictx = (struct htp_im2col_context *) data; \ @@ -143,21 +316,27 @@ IM2COL_PATCHEMBED_BODY(im2col_patchembed_f32_thread, float, hvx_copy_f32_uu, hvx struct htp_thread_trace * restrict tr = &octx->ctx->trace[ith]; \ const struct htp_tensor * restrict src1 = octx->src[1]; \ const struct htp_tensor * restrict dst = octx->dst; \ - const uint32_t N = src1->ne[3], IC = src1->ne[2], IH = src1->ne[1], IW = src1->ne[0]; \ - const uint32_t KH = octx->src[0]->ne[1], KW = octx->src[0]->ne[0]; \ - const uint32_t OH = dst->ne[2], OW = dst->ne[1]; \ - const uint32_t patch_stride = IC * KH * KW; \ - const float * restrict src_data = (const float *) src1->data; \ - DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \ - dma_queue * dmaq = octx->ctx->dma[ith]; \ - uint8_t * src_base = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread; \ - uint8_t * dst_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \ - float * srcb = (float *) src_base; \ - DST_CTYPE * dstb = (DST_CTYPE *) dst_base; \ - const uint32_t row_end_max = ictx->pe_row_base + ictx->pe_nrows; \ - const uint32_t per_thread = ictx->pe_rows_per_thread; \ - const uint32_t row_start = ictx->pe_row_base + per_thread * ith; \ - const uint32_t row_end = MIN(row_start + per_thread, row_end_max); \ + const int32_t is_2D = octx->op_params[6] == 1; \ + const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; \ + const uint32_t IC = is_2D ? src1->ne[2] : src1->ne[1]; \ + const uint32_t IH = is_2D ? src1->ne[1] : 1; \ + const uint32_t IW = src1->ne[0]; \ + const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; \ + const uint32_t KW = octx->src[0]->ne[0]; \ + const uint32_t OH = is_2D ? dst->ne[2] : 1; \ + const uint32_t OW = dst->ne[1]; \ + const uint32_t patch_stride = IC * KH * KW; \ + const float * restrict src_data = (const float *) src1->data; \ + DST_CTYPE * restrict dst_data = (DST_CTYPE *) dst->data; \ + dma_queue * dmaq = octx->ctx->dma[ith]; \ + uint8_t * src_base = ictx->pe_vtcm_src + ith * ictx->pe_src_size_per_thread; \ + uint8_t * dst_base = ictx->pe_vtcm_dst + ith * ictx->pe_dst_size_per_thread; \ + float * srcb = (float *) src_base; \ + DST_CTYPE * dstb = (DST_CTYPE *) dst_base; \ + const uint32_t row_end_max = ictx->pe_row_base + ictx->pe_nrows; \ + const uint32_t per_thread = ictx->pe_rows_per_thread; \ + const uint32_t row_start = ictx->pe_row_base + per_thread * ith; \ + const uint32_t row_end = MIN(row_start + per_thread, row_end_max); \ if (row_start >= row_end) \ return; \ for (uint32_t r = row_start; r < row_end; r++) { \ @@ -209,21 +388,30 @@ static bool im2col_use_patchembed_dma(const struct htp_ops_context * octx) { const int32_t p0 = octx->op_params[2], p1 = octx->op_params[3]; const int32_t d0 = octx->op_params[4], d1 = octx->op_params[5]; const int is_2D = octx->op_params[6] == 1; - if (!is_2D) { - return false; - } if (octx->dst->type != HTP_TYPE_F16 && octx->dst->type != HTP_TYPE_F32) { return false; } - const uint32_t KH = octx->src[0]->ne[1], KW = octx->src[0]->ne[0]; - if (s0 != (int32_t) KW || s1 != (int32_t) KH) { - return false; // non-overlapping + const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; + const uint32_t KW = octx->src[0]->ne[0]; + if (s0 != (int32_t) KW) { + return false; // non-overlapping (width) } - if (p0 != 0 || p1 != 0) { - return false; // no padding + if (p0 != 0) { + return false; // no padding (width) } - if (d0 != 1 || d1 != 1) { - return false; // no dilation + if (d0 != 1) { + return false; // no dilation (width) + } + if (is_2D) { + if (s1 != (int32_t) KH) { + return false; // non-overlapping (height) + } + if (p1 != 0) { + return false; // no padding (height) + } + if (d1 != 1) { + return false; // no dilation (height) + } } return true; } @@ -233,8 +421,11 @@ static bool im2col_use_patchembed_dma(const struct htp_ops_context * octx) { static bool im2col_patchembed_dma_fits(struct htp_ops_context * octx, struct htp_im2col_context * ictx, uint32_t n_threads) { - const uint32_t IC = octx->src[1]->ne[2], IW = octx->src[1]->ne[0]; - const uint32_t KH = octx->src[0]->ne[1], KW = octx->src[0]->ne[0]; + const int32_t is_2D = octx->op_params[6] == 1; + const uint32_t IC = is_2D ? octx->src[1]->ne[2] : octx->src[1]->ne[1]; + const uint32_t IW = octx->src[1]->ne[0]; + const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; + const uint32_t KW = octx->src[0]->ne[0]; const uint32_t OW = octx->dst->ne[1]; const uint32_t patch_stride = IC * KH * KW; @@ -257,6 +448,45 @@ static bool im2col_patchembed_dma_fits(struct htp_ops_context * octx, return true; } +// Sizes a per-thread 2x(src,dst) VTCM ping-pong for the blocked general kernel. +// Stages Wb=(owb-1)*s0+(KW-1)*d0+1 source cols per (iic,ikh) row and owb patches +// of dst. Picks the largest owb that fits; returns false if even owb=1 does not. +static bool im2col_blocked_dma_fits(struct htp_ops_context * octx, + struct htp_im2col_context * ictx, + uint32_t n_threads) { + const int32_t is_2D = octx->op_params[6] == 1; + const int32_t s0 = octx->op_params[0]; + const int32_t d0 = octx->op_params[4]; + const uint32_t IC = is_2D ? octx->src[1]->ne[2] : octx->src[1]->ne[1]; + const uint32_t KH = is_2D ? octx->src[0]->ne[1] : 1; + const uint32_t KW = octx->src[0]->ne[0]; + const uint32_t OW = octx->dst->ne[1]; + const uint32_t patch_stride = IC * KH * KW; + const uint32_t dst_elem = (octx->dst->type == HTP_TYPE_F16) ? sizeof(__fp16) : sizeof(float); + + for (uint32_t owb = (OW < 256 ? OW : 256); owb >= 1; owb--) { + const uint32_t Wb = (owb - 1) * (uint32_t) s0 + (KW - 1) * (uint32_t) d0 + 1; + const uint32_t src_row_bytes = hex_round_up(IC * KH * Wb * sizeof(float), 256); + const uint32_t dst_row_bytes = hex_round_up(owb * patch_stride * dst_elem, 256); + struct htp_im2col_vtcm_layout L; + htp_im2col_vtcm_layout_build(&L, src_row_bytes, dst_row_bytes, n_threads); + if (L.total_bytes <= octx->ctx->vtcm_size) { + uint8_t * const base = octx->ctx->vtcm_base; + ictx->pe_owb = owb; + ictx->pe_wb = Wb; + ictx->pe_src_row_bytes = src_row_bytes; + ictx->pe_dst_row_bytes = dst_row_bytes; + ictx->pe_vtcm_src = VTCM_LAYOUT_PTR(uint8_t, base, L.off_src); + ictx->pe_vtcm_dst = VTCM_LAYOUT_PTR(uint8_t, base, L.off_dst); + ictx->pe_src_size_per_thread = (uint32_t) L.src_bytes_per_thread; + ictx->pe_dst_size_per_thread = (uint32_t) L.dst_bytes_per_thread; + return true; + } + if (owb == 1) break; // avoid unsigned underflow + } + return false; +} + int op_im2col(struct htp_ops_context * octx) { const struct htp_tensor * src1 = octx->src[1]; const struct htp_tensor * dst = octx->dst; @@ -270,8 +500,9 @@ int op_im2col(struct htp_ops_context * octx) { return HTP_STATUS_OK; } - const uint32_t N = src1->ne[3]; - const uint32_t OH = dst->ne[2]; + const int32_t is_2D = octx->op_params[6] == 1; + const uint32_t N = is_2D ? src1->ne[3] : src1->ne[2]; + const uint32_t OH = is_2D ? dst->ne[2] : 1; const uint32_t OW = dst->ne[1]; const uint32_t total_patches = N * OH * OW; const uint32_t total_rows = N * OH; @@ -280,8 +511,11 @@ int op_im2col(struct htp_ops_context * octx) { uint32_t npatches = total_patches; if (octx->ctx->mdev.count > 1) { const uint32_t patch_size = dst->nb[1]; - const uint32_t patches_per_chunk = (patch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(patch_size, HEX_L2_LINE_SIZE)) : 1; - const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_patches, htp_tensor_mdev_data_aligned(dst) ? patches_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); + const uint32_t patches_per_chunk = + (patch_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(patch_size, HEX_L2_LINE_SIZE)) : 1; + const struct htp_tensor_mdev_range range = + htp_tensor_mdev_partition(total_patches, htp_tensor_mdev_data_aligned(dst) ? patches_per_chunk : 0, + octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); patch_base = range.start; npatches = range.count; } @@ -290,8 +524,11 @@ int op_im2col(struct htp_ops_context * octx) { uint32_t nrows = total_rows; if (octx->ctx->mdev.count > 1) { const uint32_t row_size = dst->nb[2]; - const uint32_t rows_per_chunk = (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1; - const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); + const uint32_t rows_per_chunk = + (row_size > 0) ? (HEX_L2_LINE_SIZE / hex_gcd_u32(row_size, HEX_L2_LINE_SIZE)) : 1; + const struct htp_tensor_mdev_range range = + htp_tensor_mdev_partition(total_rows, htp_tensor_mdev_data_aligned(dst) ? rows_per_chunk : 0, + octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); row_base = range.start; nrows = range.count; } @@ -309,22 +546,32 @@ int op_im2col(struct htp_ops_context * octx) { ictx.npatches_per_thread = (npatches + n_threads - 1) / n_threads; // Clean non-overlapping patch-embed -> DMA kernel (if it fits VTCM); - // everything else (padding/dilation/stride edges) -> pure-DDR kernel. - if (im2col_use_patchembed_dma(octx) && nrows > 0) { + // everything else (padding/dilation/stride edges) -> blocked-staging DMA + // kernel; if neither fits VTCM -> pure-DDR kernel. + if (nrows > 0) { const uint32_t pth = MIN(octx->n_threads, nrows); - if (pth > 0 && im2col_patchembed_dma_fits(octx, &ictx, pth)) { - ictx.pe_row_base = row_base; - ictx.pe_nrows = nrows; - ictx.pe_rows_per_thread = (nrows + pth - 1) / pth; - if (dst->type == HTP_TYPE_F16) { - work_queue_run(octx->ctx->work_queue, im2col_patchembed_dma_thread, &ictx, pth); - } else { - work_queue_run(octx->ctx->work_queue, im2col_patchembed_dma_f32_thread, &ictx, pth); + if (pth > 0) { + ictx.pe_row_base = row_base; + ictx.pe_nrows = nrows; + const bool exact = im2col_use_patchembed_dma(octx); + if (exact && im2col_patchembed_dma_fits(octx, &ictx, pth)) { + ictx.pe_rows_per_thread = (nrows + pth - 1) / pth; + work_queue_run(octx->ctx->work_queue, + dst->type == HTP_TYPE_F16 ? im2col_patchembed_dma_thread + : im2col_patchembed_dma_f32_thread, &ictx, pth); + return HTP_STATUS_OK; + } + if (!exact && im2col_blocked_dma_fits(octx, &ictx, pth)) { + ictx.pe_rows_per_thread = (nrows + pth - 1) / pth; + work_queue_run(octx->ctx->work_queue, + dst->type == HTP_TYPE_F16 ? im2col_blocked_dma_thread + : im2col_blocked_dma_f32_thread, &ictx, pth); + return HTP_STATUS_OK; } - return HTP_STATUS_OK; } - // else: doesn't fit -> fall through to the pure-DDR kernel below. } + // Fall through to pure-DDR. + if (npatches == 0) { return HTP_STATUS_OK; From 2b1847030cef76ef315eaee0b7ae0cdcd4fb15ff Mon Sep 17 00:00:00 2001 From: Todor Boinovski Date: Fri, 18 Sep 2026 15:05:10 -0700 Subject: [PATCH 4/9] hexagon: add ROLL op support (#29105) --- ggml/src/ggml-hexagon/ggml-hexagon.cpp | 30 +++ ggml/src/ggml-hexagon/htp/CMakeLists.txt | 1 + ggml/src/ggml-hexagon/htp/htp-ctx.h | 1 + ggml/src/ggml-hexagon/htp/htp-ops.h | 1 + ggml/src/ggml-hexagon/htp/main.c | 3 + ggml/src/ggml-hexagon/htp/roll-ops.c | 316 +++++++++++++++++++++++ 6 files changed, 352 insertions(+) create mode 100644 ggml/src/ggml-hexagon/htp/roll-ops.c diff --git a/ggml/src/ggml-hexagon/ggml-hexagon.cpp b/ggml/src/ggml-hexagon/ggml-hexagon.cpp index f6f2fdd28..766d1234f 100644 --- a/ggml/src/ggml-hexagon/ggml-hexagon.cpp +++ b/ggml/src/ggml-hexagon/ggml-hexagon.cpp @@ -5699,6 +5699,7 @@ static htp_op_code op_remap_to_htp(const ggml_tensor * t) { case GGML_OP_TRI: return HTP_OP_TRI; case GGML_OP_PAD: return HTP_OP_PAD; case GGML_OP_IM2COL: return HTP_OP_IM2COL; + case GGML_OP_ROLL: return HTP_OP_ROLL; case GGML_OP_UNARY: switch (ggml_get_unary_op(t)) { @@ -6631,6 +6632,31 @@ static bool ggml_hexagon_supported_fill(const struct ggml_hexagon_session * sess GGML_UNUSED(sess); } +static bool ggml_hexagon_supported_roll(const struct ggml_hexagon_session * sess, const struct ggml_tensor * op) { + GGML_UNUSED(sess); + + const struct ggml_tensor * src0 = op->src[0]; + const struct ggml_tensor * dst = op; + + if (src0->type != GGML_TYPE_F32 || dst->type != GGML_TYPE_F32) { + return false; + } + + if (!ggml_are_same_shape(src0, dst)) { + return false; + } + + if (src0->nb[0] != ggml_type_size(src0->type) || dst->nb[0] != ggml_type_size(dst->type)) { + return false; + } + + if (!ggml_is_contiguous(dst)) { + return false; + } + + return true; +} + static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, const struct ggml_tensor * op) { auto dev_ctx = static_cast(dev->context); auto sess = dev_ctx->session(); @@ -6798,6 +6824,10 @@ static bool ggml_backend_hexagon_device_supports_op(ggml_backend_dev_t dev, cons supp = ggml_hexagon_supported_pad(sess, op); break; + case GGML_OP_ROLL: + supp = ggml_hexagon_supported_roll(sess, op); + break; + default: break; } diff --git a/ggml/src/ggml-hexagon/htp/CMakeLists.txt b/ggml/src/ggml-hexagon/htp/CMakeLists.txt index 77f3ee39d..821f08c0b 100644 --- a/ggml/src/ggml-hexagon/htp/CMakeLists.txt +++ b/ggml/src/ggml-hexagon/htp/CMakeLists.txt @@ -43,6 +43,7 @@ add_library(${HTP_LIB} SHARED pad-ops.c argsort-ops.c im2col-ops.c + roll-ops.c allreduce-ops.c ) diff --git a/ggml/src/ggml-hexagon/htp/htp-ctx.h b/ggml/src/ggml-hexagon/htp/htp-ctx.h index 3b60c8bdb..cfb46a9ca 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ctx.h +++ b/ggml/src/ggml-hexagon/htp/htp-ctx.h @@ -175,5 +175,6 @@ int op_gated_delta_net(struct htp_ops_context * octx); int op_pad(struct htp_ops_context * octx); int op_im2col(struct htp_ops_context * octx); int op_allreduce(struct htp_ops_context * octx); +int op_roll(struct htp_ops_context * octx); #endif /* HTP_CTX_H */ diff --git a/ggml/src/ggml-hexagon/htp/htp-ops.h b/ggml/src/ggml-hexagon/htp/htp-ops.h index 98a5f6d5c..65533cbc4 100644 --- a/ggml/src/ggml-hexagon/htp/htp-ops.h +++ b/ggml/src/ggml-hexagon/htp/htp-ops.h @@ -104,6 +104,7 @@ enum htp_op_code { HTP_OP_ALLREDUCE_ADD, HTP_OP_GLU_SWIGLU_CLAMP, HTP_OP_MDEV_GROUP, + HTP_OP_ROLL, HTP_OP_INVALID }; diff --git a/ggml/src/ggml-hexagon/htp/main.c b/ggml/src/ggml-hexagon/htp/main.c index 1d291e16b..4fad5de6f 100644 --- a/ggml/src/ggml-hexagon/htp/main.c +++ b/ggml/src/ggml-hexagon/htp/main.c @@ -879,6 +879,9 @@ static int execute_op(struct htp_ops_context * octx) { case HTP_OP_IM2COL: return op_im2col(octx); + case HTP_OP_ROLL: + return op_roll(octx); + case HTP_OP_CONCAT: return op_concat(octx); diff --git a/ggml/src/ggml-hexagon/htp/roll-ops.c b/ggml/src/ggml-hexagon/htp/roll-ops.c new file mode 100644 index 000000000..6faf2ac47 --- /dev/null +++ b/ggml/src/ggml-hexagon/htp/roll-ops.c @@ -0,0 +1,316 @@ +#pragma clang diagnostic ignored "-Wunused-variable" +#pragma clang diagnostic ignored "-Wunused-function" +#pragma clang diagnostic ignored "-Wunused-but-set-variable" + +#include +#include + +#include + +#include "dma-queue.h" +#include "hvx-utils.h" + +#define GGML_COMMON_DECL_C +#include "ggml-common.h" +#include "htp-ctx.h" +#include "hex-common.h" +#include "hex-profile.h" +#include "htp-ops.h" +#include "htp-tensor.h" + +struct htp_roll_context { + struct htp_ops_context * octx; + + uint32_t row_start; + uint32_t nrows; + uint32_t nrows_per_thread; + + struct fastdiv_values div_ne1; + struct fastdiv_values div_ne2_ne1; +}; + +static inline uint32_t htp_roll_wrap(int32_t i, uint32_t ne) { + if (i < 0) { + return (uint32_t) (i + (int32_t) ne); + } + if ((uint32_t) i >= ne) { + return (uint32_t) i - ne; + } + return (uint32_t) i; +} + +#define htp_roll_preamble \ + const struct htp_tensor * src0 = octx->src[0]; \ + const struct htp_tensor * dst = octx->dst; \ + \ + const uint32_t ne0 = dst->ne[0]; \ + const uint32_t ne1 = dst->ne[1]; \ + const uint32_t ne2 = dst->ne[2]; \ + const uint32_t ne3 = dst->ne[3]; \ + \ + const uint32_t nb01 = src0->nb[1]; \ + const uint32_t nb02 = src0->nb[2]; \ + const uint32_t nb03 = src0->nb[3]; \ + \ + const uint32_t nb1 = dst->nb[1]; \ + const uint32_t nb2 = dst->nb[2]; \ + const uint32_t nb3 = dst->nb[3]; \ + \ + const int32_t s0 = octx->op_params[0]; \ + const int32_t s1 = octx->op_params[1]; \ + const int32_t s2 = octx->op_params[2]; \ + const int32_t s3 = octx->op_params[3]; \ + \ + const uint32_t i0_src0 = htp_roll_wrap(-s0, ne0); \ + const uint32_t n0 = ne0 - i0_src0; + +#define htp_roll_dma_preamble dma_queue * q = octx->ctx->dma[0]; + +static inline void roll_dma_push(dma_queue * q, + uintptr_t dst, + uintptr_t src, + uint32_t dst_stride, + uint32_t src_stride, + uint32_t bytes, + uint32_t nrows) { + if (bytes == 0 || nrows == 0) { + return; + } + + if (!dma_queue_push(q, dma_make_ptr((void *) dst, (const void *) src), dst_stride, src_stride, bytes, nrows)) { + dma_queue_flush(q); + dma_queue_push(q, dma_make_ptr((void *) dst, (const void *) src), + dst_stride, src_stride, bytes, nrows); + } +} + +static inline void roll_dma_push_rows(dma_queue * q, + const struct htp_tensor * dst, + const struct htp_tensor * src0, + uint32_t dst_row, + uint32_t src_row, + uint32_t nrows, + uint32_t row_size, + uint32_t i0_src0) { + const uintptr_t dst_base = dst->data + (uintptr_t) dst_row * row_size; + const uintptr_t src_base = src0->data + (uintptr_t) src_row * row_size; + const uint32_t n0 = src0->ne[0] - i0_src0; + + roll_dma_push(q, dst_base, src_base + (uintptr_t) i0_src0 * sizeof(float), + row_size, row_size, n0 * sizeof(float), nrows); + roll_dma_push(q, dst_base + (uintptr_t) n0 * sizeof(float), src_base, + row_size, row_size, i0_src0 * sizeof(float), nrows); +} + +// Same row-wrap split as roll_dma_push_rows, but addressed with explicit byte strides so it +// also works for a src0 that is row-contiguous only (e.g. a permuted view) rather than fully packed. +static inline void roll_dma_push_range(dma_queue * q, + uintptr_t dst_row, + uintptr_t src_row, + uint32_t dst_stride, + uint32_t src_stride, + uint32_t nrows, + uint32_t i0_src0, + uint32_t n0) { + roll_dma_push(q, dst_row, src_row + (uintptr_t) i0_src0 * sizeof(float), + dst_stride, src_stride, n0 * sizeof(float), nrows); + roll_dma_push(q, dst_row + (uintptr_t) n0 * sizeof(float), src_row, + dst_stride, src_stride, i0_src0 * sizeof(float), nrows); +} + +static int roll_dma_f32_contiguous(struct htp_ops_context * octx) { + htp_roll_preamble; + htp_roll_dma_preamble; + + const uint32_t row_size = ne0 * sizeof(float); + + if (s1 == 0 && s2 == 0 && s3 == 0) { + roll_dma_push_rows(q, dst, src0, 0, 0, ne1 * ne2 * ne3, row_size, i0_src0); + dma_queue_flush(q); + return HTP_STATUS_OK; + } + + if (s1 == 0) { + const uint32_t i2_src0 = htp_roll_wrap(-s2, ne2); + for (uint32_t i3 = 0; i3 < ne3; i3++) { + const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3); + const uint32_t dst_row0 = i3 * ne2 * ne1; + const uint32_t src_row0 = (i03 * ne2 + i2_src0) * ne1; + const uint32_t n2_first = ne2 - i2_src0; + + roll_dma_push_rows(q, dst, src0, dst_row0, src_row0, n2_first * ne1, + row_size, i0_src0); + roll_dma_push_rows(q, dst, src0, dst_row0 + n2_first * ne1, i03 * ne2 * ne1, + i2_src0 * ne1, row_size, i0_src0); + } + + dma_queue_flush(q); + return HTP_STATUS_OK; + } + + const uint32_t i1_src0 = htp_roll_wrap(-s1, ne1); + const uint32_t n1_first = ne1 - i1_src0; + + for (uint32_t i3 = 0; i3 < ne3; i3++) { + const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3); + for (uint32_t i2 = 0; i2 < ne2; i2++) { + const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2); + const uint32_t dst_row0 = (i3 * ne2 + i2) * ne1; + const uint32_t src_row0 = (i03 * ne2 + i02) * ne1; + + roll_dma_push_rows(q, dst, src0, dst_row0, src_row0 + i1_src0, + n1_first, row_size, i0_src0); + roll_dma_push_rows(q, dst, src0, dst_row0 + n1_first, src_row0, + i1_src0, row_size, i0_src0); + } + } + + dma_queue_flush(q); + return HTP_STATUS_OK; +} + +// DMA path for a row-contiguous but otherwise arbitrarily strided src0 (e.g. a permuted view). +// Same row-wrap split as above, one DMA push per (i2,i3), addressed via the real nb01/nb02/nb03 +// instead of assuming a packed layout. +static int roll_dma_f32_strided(struct htp_ops_context * octx) { + htp_roll_preamble; + htp_roll_dma_preamble; + + const uint32_t i1_src0 = htp_roll_wrap(-s1, ne1); + const uint32_t n1_first = ne1 - i1_src0; + + for (uint32_t i3 = 0; i3 < ne3; i3++) { + const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3); + for (uint32_t i2 = 0; i2 < ne2; i2++) { + const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2); + + const uintptr_t dst_row0 = dst->data + (uintptr_t) i2 * nb2 + (uintptr_t) i3 * nb3; + const uintptr_t src_row0 = src0->data + (uintptr_t) i02 * nb02 + (uintptr_t) i03 * nb03; + + roll_dma_push_range(q, dst_row0, src_row0 + (uintptr_t) i1_src0 * nb01, + nb1, nb01, n1_first, i0_src0, n0); + roll_dma_push_range(q, dst_row0 + (uintptr_t) n1_first * nb1, src_row0, + nb1, nb01, i1_src0, i0_src0, n0); + } + } + + dma_queue_flush(q); + return HTP_STATUS_OK; +} + +static void roll_thread_f32(unsigned int nth, unsigned int ith, void * data) { + struct htp_roll_context * rctx = (struct htp_roll_context *) data; + struct htp_ops_context * octx = rctx->octx; + + htp_roll_preamble; + + const uint32_t row_start = rctx->row_start + rctx->nrows_per_thread * ith; + const uint32_t row_end = MIN(row_start + rctx->nrows_per_thread, rctx->row_start + rctx->nrows); + if (row_start >= row_end) { + return; + } + + struct htp_thread_trace * tr = &octx->ctx->trace[ith]; + htp_trace_event_start(tr, HTP_TRACE_EVT_HVX_COMP, row_start); + + for (uint32_t row = row_start; row < row_end; row++) { + const uint32_t i3 = fastdiv(row, &rctx->div_ne2_ne1); + const uint32_t rem = row - i3 * ne2 * ne1; + const uint32_t i2 = fastdiv(rem, &rctx->div_ne1); + const uint32_t i1 = rem - i2 * ne1; + + const uint32_t i01 = htp_roll_wrap((int32_t) i1 - s1, ne1); + const uint32_t i02 = htp_roll_wrap((int32_t) i2 - s2, ne2); + const uint32_t i03 = htp_roll_wrap((int32_t) i3 - s3, ne3); + + const uint8_t * src_row = (const uint8_t *) src0->data + i01*nb01 + i02*nb02 + i03*nb03; + uint8_t * dst_row = (uint8_t *) dst->data + i1*nb1 + i2*nb2 + i3*nb3; + + hex_l2fetch(src_row + i0_src0 * sizeof(float), n0 * sizeof(float), ne0 * sizeof(float), 1); + hvx_copy_uu(dst_row, src_row + i0_src0 * sizeof(float), n0, sizeof(float)); + + if (i0_src0 != 0) { + hex_l2fetch(src_row, i0_src0 * sizeof(float), ne0 * sizeof(float), 1); + hvx_copy_uu(dst_row + n0 * sizeof(float), src_row, i0_src0, sizeof(float)); + } + } + + htp_trace_event_stop(tr, HTP_TRACE_EVT_HVX_COMP, row_start); + + FARF(HIGH, "roll %d/%d: (%ux%ux%ux%u) rows %u:%u shift=(%d,%d,%d,%d)\n", + ith, nth, ne0, ne1, ne2, ne3, + row_start, row_end, s0, s1, s2, s3); +} + +int execute_op_roll_f32(struct htp_ops_context * octx) { + htp_roll_preamble; + + if (src0->type != HTP_TYPE_F32 || dst->type != HTP_TYPE_F32) { + FARF(ERROR, "roll: unsupported type %u -> %u\n", src0->type, dst->type); + return HTP_STATUS_NO_SUPPORT; + } + + if (src0->nb[0] != sizeof(float) || dst->nb[0] != sizeof(float)) { + FARF(ERROR, "roll: unsupported nb0 %u -> %u\n", src0->nb[0], dst->nb[0]); + return HTP_STATUS_NO_SUPPORT; + } + + if (src0->ne[0] != ne0 || src0->ne[1] != ne1 || + src0->ne[2] != ne2 || src0->ne[3] != ne3) { + FARF(ERROR, "roll: shape mismatch\n"); + return HTP_STATUS_INVAL_PARAMS; + } + + if (octx->flags & HTP_OPFLAGS_SKIP_COMPUTE) { + return HTP_STATUS_OK; + } + + const uint32_t total_rows = ne1 * ne2 * ne3; + const size_t dst_row_size = ne0 * sizeof(float); + + uint32_t row_start = 0; + uint32_t nrows = total_rows; + + if (octx->ctx->mdev.count > 1) { + uint32_t rows_per_chunk = 0; + htp_tensor_mdev_rows_per_chunk(dst, sizeof(float), (uint32_t) dst_row_size, &rows_per_chunk); + const struct htp_tensor_mdev_range range = htp_tensor_mdev_partition(total_rows, rows_per_chunk, octx->ctx->mdev.idx, octx->ctx->mdev.count, &octx->ctx->mdev.count_div); + row_start = range.start; + nrows = range.count; + } + + if (nrows == 0) { + return HTP_STATUS_OK; + } + + if (octx->ctx->mdev.count <= 1) { + if (htp_tensor_is_contiguous(src0, sizeof(float)) && htp_tensor_is_contiguous(dst, sizeof(float))) { + return roll_dma_f32_contiguous(octx); + } + return roll_dma_f32_strided(octx); + } + + const uint32_t n_threads = octx->n_threads; + struct htp_roll_context rctx = { + .octx = octx, + .row_start = row_start, + .nrows = nrows, + .nrows_per_thread = fastdiv(nrows + n_threads - 1, &octx->n_threads_div), + .div_ne1 = init_fastdiv_values(dst->ne[1]), + .div_ne2_ne1 = init_fastdiv_values(dst->ne[2] * dst->ne[1]), + }; + + work_queue_run(octx->ctx->work_queue, roll_thread_f32, &rctx, n_threads); + + return HTP_STATUS_OK; +} + +int op_roll(struct htp_ops_context * octx) { + switch (octx->src[0]->type) { + case HTP_TYPE_F32: + return execute_op_roll_f32(octx); + + default: + return HTP_STATUS_NO_SUPPORT; + } +} From 60081bb2b5b3294165a4d67c5cbeebe74c868014 Mon Sep 17 00:00:00 2001 From: dsproule Date: Fri, 18 Sep 2026 16:32:31 -0700 Subject: [PATCH 5/9] opencl: add support for bin kernel `flash_attn_f32_f16_bin` (#29046) * opencl: add `flash_attn_f32_f16_bin` * opencl: guarded prefill fa --- ggml/src/ggml-opencl/CMakeLists.txt | 1 + ggml/src/ggml-opencl/ggml-opencl.cpp | 483 ++++++++++++++++++ .../ggml-opencl/kernels/flash_attn_repack.cl | 92 ++++ 3 files changed, 576 insertions(+) create mode 100644 ggml/src/ggml-opencl/kernels/flash_attn_repack.cl diff --git a/ggml/src/ggml-opencl/CMakeLists.txt b/ggml/src/ggml-opencl/CMakeLists.txt index 53e938618..ff5e8ef46 100644 --- a/ggml/src/ggml-opencl/CMakeLists.txt +++ b/ggml/src/ggml-opencl/CMakeLists.txt @@ -233,6 +233,7 @@ set(GGML_OPENCL_KERNELS mul_mm_f16_f32_kq_kqv conv2d conv2d_f16_f32 + flash_attn_repack flash_attn_pre_f16 flash_attn_f32_f16 flash_attn_f32_q8_0 diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index 1c26797b9..fe7377b2d 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -567,6 +567,16 @@ struct ggml_opencl_fa_kernels { // attempted (variant, (dk, dv)) // all attempted FA kernels appear here, but those not registered failed compilation std::set>> variant_attempted; + + // FA bin kernels +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + cl_kernel kernel_flash_attn_f32_f16_bin; + + cl_kernel kernel_repack_q_for_wmm; + cl_kernel kernel_repack_k_for_wmm; + cl_kernel kernel_repack_v_for_wmm; + cl_kernel kernel_repack_mask_for_wmm; +#endif }; #ifdef GGML_OPENCL_USE_ADRENO_KERNELS @@ -5172,6 +5182,43 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) { CL_CHECK(clReleaseProgram(prog)); GGML_LOG_CONT("."); } + + // repack + { +#ifdef GGML_OPENCL_EMBED_KERNELS + const std::string kernel_src { + #include "flash_attn_repack.cl.h" + }; +#else + const std::string kernel_src = read_file("flash_attn_repack.cl"); +#endif + cl_program prog = + build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts); + + CL_CHECK((backend_ctx->fa.kernel_repack_q_for_wmm = clCreateKernel(prog, "kernel_repack_q_for_wmm", &err), err)); + CL_CHECK((backend_ctx->fa.kernel_repack_k_for_wmm = clCreateKernel(prog, "kernel_repack_k_for_wmm", &err), err)); + CL_CHECK((backend_ctx->fa.kernel_repack_v_for_wmm = clCreateKernel(prog, "kernel_repack_v_for_wmm", &err), err)); + CL_CHECK((backend_ctx->fa.kernel_repack_mask_for_wmm = clCreateKernel(prog, "kernel_repack_mask_for_wmm", &err), err)); + GGML_LOG_CONT("."); + } + + // kernel_flash_attn_f32_f16_bin + { + size_t bin_size = 0; + backend_ctx->fa.kernel_flash_attn_f32_f16_bin = nullptr; + + if (use_adreno_bin_kernels(backend_ctx)) { + const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("flash_attn_f32_f16_wmm", &bin_size); + if (kernel_bin && bin_size > 0) { + cl_program prog = + build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, CL_moe_compile_opts, bin_size); + + CL_CHECK((backend_ctx->fa.kernel_flash_attn_f32_f16_bin = clCreateKernel(prog, "flash_attn_f32_f16", &err), err)); + CL_CHECK(clReleaseProgram(prog)); + GGML_LOG_CONT("."); + } + } + } #endif // GGML_OPENCL_USE_ADRENO_KERNELS GGML_LOG_CONT("\n"); backend_ctx->kernels_loaded = true; @@ -8532,6 +8579,28 @@ inline bool use_q4_0_bin_kernels(const ggml_backend_opencl_context *backend_ctx, #endif } +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS +static bool use_fa_bin_kernels_prefill(const ggml_backend_opencl_context * backend_ctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v) { + if (backend_ctx->fa.kernel_flash_attn_f32_f16_bin == nullptr) { + return false; + } + + const bool is_mixed = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16; + const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0; + + const int n_q = q->ne[1]; + const int dk = q->ne[0]; + const int dv = v->ne[0]; + + constexpr bool prefill_only = true; + + return (backend_ctx->gpu_family == GPU_FAMILY::ADRENO && + (is_mixed || is_q8_0) && (dk == dv) + && (dk == 64 || dk == 128 || dk == 256 || dk == 512) + && (!prefill_only || n_q != 1)); +} +#endif + // The flat-GEMV large-m escape is OPT-IN (GGML_OPENCL_FLAT_LARGE_M=1) because it // is SLOWER than the route it replaces, not because it is unsafe. It was first // parked on the theory that it out-of-bounds-writes at vocab-scale shapes; that @@ -8990,6 +9059,11 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te case GGML_OP_MEAN: return op->src[0]->type == GGML_TYPE_F32; case GGML_OP_FLASH_ATTN_EXT: { +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + if (use_fa_bin_kernels_prefill(backend_ctx, op->src[0], op->src[1], op->src[2])) { + return true; + } +#endif // The E17 compilers segfault while building FA kernels, skip E17 for now if (adreno_e17_compiler_quirks(backend_ctx)) { return false; @@ -17198,6 +17272,407 @@ static void ggml_cl_adreno_xmem_attn_run( #endif // GGML_OPENCL_USE_ADRENO_KERNELS +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS +static void ggml_cl_flash_attn_prefill_bin(ggml_backend_t backend, const ggml_tensor * q, const ggml_tensor * k, ggml_tensor * dst) { + const ggml_tensor * v = dst->src[2]; + const ggml_tensor * mask = dst->src[3]; + const ggml_tensor * sinks = dst->src[4]; + GGML_ASSERT(q->extra); + GGML_ASSERT(k->extra); + GGML_ASSERT(v->extra); + GGML_ASSERT(dst->extra); + if (mask) { + GGML_ASSERT(mask->extra); + } + if (sinks) { + GGML_ASSERT(sinks->extra); + } + + ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context; + cl_context context = backend_ctx->context; + + const int n_q = q->ne[1]; + const int n_kv = k->ne[1]; + const int d_head_q = q->ne[0]; + const int d_head_v = v->ne[0]; + const int n_head = q->ne[2]; + const int n_head_kv = k->ne[2]; + const int n_batch = q->ne[3]; + + const std::pair dk_dv = {d_head_q, d_head_v}; + cl_kernel kernel = backend_ctx->fa.kernel_flash_attn_f32_f16_bin; + GGML_ASSERT(kernel != NULL); + + ggml_tensor_extra_cl * extra_q = (ggml_tensor_extra_cl *)q->extra; + ggml_tensor_extra_cl * extra_k = (ggml_tensor_extra_cl *)k->extra; + ggml_tensor_extra_cl * extra_v = (ggml_tensor_extra_cl *)v->extra; + ggml_tensor_extra_cl * extra_o = (ggml_tensor_extra_cl *)dst->extra; + ggml_tensor_extra_cl * extra_mask = mask ? (ggml_tensor_extra_cl *)mask->extra : NULL; + ggml_tensor_extra_cl * extra_sinks = sinks ? (ggml_tensor_extra_cl *)sinks->extra : NULL; + + cl_ulong offset_q = extra_q->offset + q->view_offs; + cl_ulong offset_o = extra_o->offset + dst->view_offs; + + cl_mem mask_buffer = extra_mask ? extra_mask->data_device : NULL; + cl_ulong offset_mask = extra_mask ? extra_mask->offset + mask->view_offs : 0; + cl_mem sinks_buffer = extra_sinks ? extra_sinks->data_device : NULL; + cl_ulong offset_sinks = extra_sinks ? extra_sinks->offset + sinks->view_offs : 0; + + const cl_ulong q_nb1 = q->nb[1]; + const cl_ulong q_nb2 = q->nb[2]; + const cl_ulong q_nb3 = q->nb[3]; + + cl_mem k_data_device = extra_k->data_device; + cl_ulong offset_k = extra_k->offset + k->view_offs; + cl_ulong k_nb1 = k->nb[1]; + cl_ulong k_nb2 = k->nb[2]; + cl_ulong k_nb3 = k->nb[3]; + + cl_mem v_data_device = extra_v->data_device; + cl_ulong offset_v = extra_v->offset + v->view_offs; + cl_ulong v_nb1 = v->nb[1]; + cl_ulong v_nb2 = v->nb[2]; + cl_ulong v_nb3 = v->nb[3]; + + const cl_ulong o_nb1 = dst->nb[1]; + const cl_ulong o_nb2 = dst->nb[2]; + const cl_ulong o_nb3 = dst->nb[3]; + + const cl_ulong mask_nb1 = mask ? mask->nb[1] : 0; + const cl_ulong mask_nb2 = mask ? mask->nb[2] : 0; + const cl_ulong mask_nb3 = mask ? mask->nb[3] : 0; + const int mask_ne2 = mask ? mask->ne[2] : 0; + const int mask_ne3 = mask ? mask->ne[3] : 0; + + float * params = (float *)dst->op_params; + float scale = params[0]; + float max_bias = params[1]; + float logit_softcap = params[2]; + + const int is_causal = (mask == NULL && n_q > 1 && n_q == n_kv); // redundant n_q > 1 check ? + + const int n_head_log2_val = n_head > 0 ? 1u << (int)floorf(log2f((float)n_head)) : 0; + const float n_head_log2_f = n_head_log2_val > 0 ? (float)n_head_log2_val : 1.0f; + const float m0 = powf(2.0f, -(max_bias) / n_head_log2_f); + const float m1 = powf(2.0f, -(max_bias / 2.0f) / n_head_log2_f); + + const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0; + + ggml_cl_flash_attn_temp_buffer temp_k; + ggml_cl_flash_attn_temp_buffer temp_v; + ggml_cl_flash_attn_temp_buffer temp_k_aos; + ggml_cl_flash_attn_temp_buffer temp_v_aos; + + if (is_q8_0) { + ggml_cl_flash_attn_reconstruct_aos( + backend_ctx, k, temp_k_aos, k_data_device, offset_k, k_nb1, k_nb2, k_nb3); + + ggml_cl_flash_attn_reconstruct_aos( + backend_ctx, v, temp_v_aos, v_data_device, offset_v, v_nb1, v_nb2, v_nb3); + + bool k_done = ggml_cl_flash_attn_dequant_kv_gpu( + backend_ctx, k, GGML_TYPE_F16, k_data_device, offset_k, k_nb1, k_nb2, k_nb3, + temp_k, k_data_device, offset_k, k_nb1, k_nb2, k_nb3); + + bool v_done = ggml_cl_flash_attn_dequant_kv_gpu( + backend_ctx, v, GGML_TYPE_F16, v_data_device, offset_v, v_nb1, v_nb2, v_nb3, + temp_v, v_data_device, offset_v, v_nb1, v_nb2, v_nb3); + + GGML_ASSERT(k_done && v_done); + } + + // Allocate input/output memory buffers + cl_mem mem_matrixQ; + cl_mem mem_matrixK; + cl_mem mem_matrixV; + cl_mem mem_matrixO; + cl_buffer_region region; + cl_int err; + + region.origin = offset_q; + region.size = ggml_nbytes(q); + mem_matrixQ = clCreateSubBuffer(extra_q->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err); + CL_CHECK(err); + + region.origin = offset_k; + region.size = is_q8_0 ? (size_t) k_nb3 * (size_t) k->ne[3] : ggml_nbytes(k); + mem_matrixK = clCreateSubBuffer(k_data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err); + CL_CHECK(err); + + region.origin = offset_v; + region.size = is_q8_0 ? (size_t) v_nb3 * (size_t) v->ne[3] : ggml_nbytes(v); + mem_matrixV = clCreateSubBuffer(v_data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err); + CL_CHECK(err); + + region.origin = offset_o; + region.size = ggml_nbytes(dst); + mem_matrixO = clCreateSubBuffer(extra_o->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err); + CL_CHECK(err); + + cl_image_format img_fmt_1d = { CL_RGBA, CL_FLOAT}; + cl_image_desc img_desc_1d; + + // use image 1d buffer used as fallback when on mask is applied + cl_mem mem_tex_mask_fallback_1dbuf; + img_fmt_1d = { CL_RGBA, CL_HALF_FLOAT}; + memset(&img_desc_1d, 0, sizeof(img_desc_1d)); + img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc_1d.image_width = 1; + img_desc_1d.buffer = mem_matrixK; + mem_tex_mask_fallback_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_1d, &img_desc_1d, NULL, &err); + CL_CHECK(err); + + cl_mem mem_tex_matrixO_1dbuf; + img_fmt_1d = { CL_RGBA, CL_FLOAT}; + memset(&img_desc_1d, 0, sizeof(img_desc_1d)); + img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc_1d.image_width = ggml_nbytes(dst) / 4 / 4; + img_desc_1d.buffer = mem_matrixO; + mem_tex_matrixO_1dbuf = clCreateImage(context, CL_MEM_WRITE_ONLY, &img_fmt_1d, &img_desc_1d, NULL, &err); + CL_CHECK(err); + + // The bin kernel requires 2d (or 3d) buffers packed for data loading/multiplication. + // These repack kernels launch across all buffers to ensure compatibility + cl_mem mem_tex_matrixMask_1dbuf = NULL; + cl_mem mem_matrixMask = NULL; + cl_mem mem_matrixMask_padded = NULL; + cl_ulong mask_nb1_padded = mask_nb1, mask_nb2_padded = mask_nb2, mask_nb3_padded = mask_nb3; + if (extra_mask) { + // allocate mem_matrixMask w/ new padded size + size_t n_kv_padded = GGML_PAD(n_kv, 4); + size_t mask_nb_padded = n_kv_padded * sizeof(cl_half) * mask->ne[1] * mask->ne[2] * mask->ne[3]; + + // apply offset and create subBuffer for mask + region.origin = offset_mask; + region.size = ggml_nbytes(mask); + mem_matrixMask = clCreateSubBuffer(extra_mask->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err); + CL_CHECK(err); + + { + // create padded mask to contain all data + mem_matrixMask_padded = clCreateBuffer(context, CL_MEM_ALLOC_HOST_PTR, mask_nb_padded, NULL, &err); + CL_CHECK(err); + + // pass extra_mask->data_device, mem_matrixMask to kernel for copying/padding + mask_nb1_padded = (cl_ulong)n_kv_padded * sizeof(cl_half); + mask_nb2_padded = mask_nb1_padded * (cl_ulong)mask->ne[1]; + mask_nb3_padded = mask_nb2_padded * (cl_ulong)mask->ne[2]; + + cl_kernel repack_mask = backend_ctx->fa.kernel_repack_mask_for_wmm; + CL_CHECK(clSetKernelArg(repack_mask, 0, sizeof(cl_mem), &mem_matrixMask)); + CL_CHECK(clSetKernelArg(repack_mask, 1, sizeof(cl_ulong), &mask_nb1)); + CL_CHECK(clSetKernelArg(repack_mask, 2, sizeof(cl_ulong), &mask_nb2)); + CL_CHECK(clSetKernelArg(repack_mask, 3, sizeof(cl_ulong), &mask_nb3)); + CL_CHECK(clSetKernelArg(repack_mask, 4, sizeof(int), &mask_ne2)); + CL_CHECK(clSetKernelArg(repack_mask, 5, sizeof(cl_mem), &mem_matrixMask_padded)); + CL_CHECK(clSetKernelArg(repack_mask, 6, sizeof(cl_ulong), &mask_nb1_padded)); + CL_CHECK(clSetKernelArg(repack_mask, 7, sizeof(cl_ulong), &mask_nb2_padded)); + CL_CHECK(clSetKernelArg(repack_mask, 8, sizeof(cl_ulong), &mask_nb3_padded)); + + size_t repack_mask_gws[3] = {(size_t)n_kv, (size_t)mask->ne[1], (size_t)mask_ne2 * (size_t)mask->ne[3]}; + backend_ctx->enqueue_ndrange_kernel(repack_mask, 3, repack_mask_gws, NULL, dst); + } + + // use image 1d buffer for matrix Mask (padded row stride) + cl_image_format img_fmt_mask_1d = { CL_RGBA, CL_HALF_FLOAT}; + cl_image_desc img_desc_mask_1d; + memset(&img_desc_mask_1d, 0, sizeof(img_desc_mask_1d)); + img_desc_mask_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc_mask_1d.image_width = mask_nb_padded / 2 / 4; + img_desc_mask_1d.buffer = mem_matrixMask_padded; + mem_tex_matrixMask_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_mask_1d, &img_desc_mask_1d, NULL, &err); + CL_CHECK(err); + } + + // WMM QK uses repacked 3D images. + // Q image: rows, heads, packed depth. + cl_image_format img_fmt_3d = { CL_RGBA, CL_HALF_FLOAT }; + cl_image_desc img_desc_3d; + + memset(&img_desc_3d, 0, sizeof(img_desc_3d)); + img_desc_3d.image_type = CL_MEM_OBJECT_IMAGE3D; + img_desc_3d.image_width = (size_t)n_q; + img_desc_3d.image_height = (size_t)n_batch * (size_t)n_head; + img_desc_3d.image_depth = (size_t)d_head_q / 4; + cl_mem img_q_wmm = NULL; + img_q_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err); + CL_CHECK(err); + + { + cl_kernel repack_q = backend_ctx->fa.kernel_repack_q_for_wmm; + CL_CHECK(clSetKernelArg(repack_q, 0, sizeof(cl_mem), &mem_matrixQ)); + CL_CHECK(clSetKernelArg(repack_q, 1, sizeof(cl_ulong), &q_nb1)); + CL_CHECK(clSetKernelArg(repack_q, 2, sizeof(cl_ulong), &q_nb2)); + CL_CHECK(clSetKernelArg(repack_q, 3, sizeof(cl_ulong), &q_nb3)); + CL_CHECK(clSetKernelArg(repack_q, 4, sizeof(int), &n_head)); + CL_CHECK(clSetKernelArg(repack_q, 5, sizeof(cl_mem), &img_q_wmm)); + + size_t repack_q_gws[3] = {(size_t)d_head_q / 4, (size_t)n_q, (size_t)n_batch * (size_t)n_head}; + backend_ctx->enqueue_ndrange_kernel(repack_q, 3, repack_q_gws, NULL, dst); + } + + // K image: columns, row groups, KV heads. + const size_t n_kv_row4 = ((size_t)n_kv + 3) / 4; + + memset(&img_desc_3d, 0, sizeof(img_desc_3d)); + img_desc_3d.image_type = CL_MEM_OBJECT_IMAGE3D; + img_desc_3d.image_width = (size_t)d_head_q; + img_desc_3d.image_height = n_kv_row4; + img_desc_3d.image_depth = (size_t)n_batch * (size_t)n_head_kv; + cl_mem img_k_wmm = NULL; + img_k_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err); + CL_CHECK(err); + + { + cl_kernel repack_k = backend_ctx->fa.kernel_repack_k_for_wmm; + CL_CHECK(clSetKernelArg(repack_k, 0, sizeof(cl_mem), &mem_matrixK)); + CL_CHECK(clSetKernelArg(repack_k, 1, sizeof(cl_ulong), &k_nb1)); + CL_CHECK(clSetKernelArg(repack_k, 2, sizeof(cl_ulong), &k_nb2)); + CL_CHECK(clSetKernelArg(repack_k, 3, sizeof(cl_ulong), &k_nb3)); + CL_CHECK(clSetKernelArg(repack_k, 4, sizeof(int), &n_head_kv)); + CL_CHECK(clSetKernelArg(repack_k, 5, sizeof(int), &n_kv)); + CL_CHECK(clSetKernelArg(repack_k, 6, sizeof(cl_mem), &img_k_wmm)); + + size_t repack_k_gws[3] = {(size_t)d_head_q, n_kv_row4, (size_t)n_batch * (size_t)n_head_kv}; + backend_ctx->enqueue_ndrange_kernel(repack_k, 3, repack_k_gws, NULL, dst); + } + + // V image: kv-rows (contracted), packed head-dim groups, KV heads. + memset(&img_desc_3d, 0, sizeof(img_desc_3d)); + img_desc_3d.image_type = CL_MEM_OBJECT_IMAGE3D; + img_desc_3d.image_width = (size_t)n_kv; + img_desc_3d.image_height = (size_t)d_head_v / 4; + img_desc_3d.image_depth = (size_t)n_batch * (size_t)n_head_kv; + cl_mem img_v_wmm = NULL; + img_v_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err); + CL_CHECK(err); + + { + cl_kernel repack_v = backend_ctx->fa.kernel_repack_v_for_wmm; + CL_CHECK(clSetKernelArg(repack_v, 0, sizeof(cl_mem), &mem_matrixV)); + CL_CHECK(clSetKernelArg(repack_v, 1, sizeof(cl_ulong), &v_nb1)); + CL_CHECK(clSetKernelArg(repack_v, 2, sizeof(cl_ulong), &v_nb2)); + CL_CHECK(clSetKernelArg(repack_v, 3, sizeof(cl_ulong), &v_nb3)); + CL_CHECK(clSetKernelArg(repack_v, 4, sizeof(int), &n_head_kv)); + CL_CHECK(clSetKernelArg(repack_v, 5, sizeof(cl_mem), &img_v_wmm)); + + size_t repack_v_gws[3] = {(size_t)d_head_v / 4, (size_t)n_kv, (size_t)n_batch * (size_t)n_head_kv}; + backend_ctx->enqueue_ndrange_kernel(repack_v, 3, repack_v_gws, NULL, dst); + } + + cl_int enable_mask = (extra_mask) ? 1 : 0; + mask_buffer = extra_mask ? mem_tex_matrixMask_1dbuf : mem_tex_mask_fallback_1dbuf; + + cl_mem mem_sinksBuf = NULL; + cl_mem mem_tex_sinks_1dbuf = NULL; + cl_int enable_sinks = (sinks_buffer != NULL) ? 1 : 0; + if (enable_sinks) { + region.origin = offset_sinks; + region.size = ggml_nbytes(sinks); + mem_sinksBuf = clCreateSubBuffer(extra_sinks->data_device, CL_MEM_READ_ONLY, CL_BUFFER_CREATE_TYPE_REGION, ®ion, &err); + CL_CHECK(err); + + cl_image_format img_fmt_sinks_1d = { CL_R, CL_FLOAT }; + cl_image_desc img_desc_sinks_1d; + memset(&img_desc_sinks_1d, 0, sizeof(img_desc_sinks_1d)); + img_desc_sinks_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc_sinks_1d.image_width = (size_t)n_head; + img_desc_sinks_1d.buffer = mem_sinksBuf; + mem_tex_sinks_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_sinks_1d, &img_desc_sinks_1d, NULL, &err); + CL_CHECK(err); + } else { + // The image obj cannot be null so we back with buffer of size 1 and use matrixK to back because it always exists + cl_image_format img_fmt_sinks_fallback = { CL_R, CL_FLOAT }; + cl_image_desc img_desc_sinks_fallback; + memset(&img_desc_sinks_fallback, 0, sizeof(img_desc_sinks_fallback)); + img_desc_sinks_fallback.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER; + img_desc_sinks_fallback.image_width = 1; + img_desc_sinks_fallback.buffer = mem_matrixK; + mem_tex_sinks_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_sinks_fallback, &img_desc_sinks_fallback, NULL, &err); + CL_CHECK(err); + } + + cl_uint arg = 0; + + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &mem_tex_matrixO_1dbuf)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &scale)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_q)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_kv)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &is_causal)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_head)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb1)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb2)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb3)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb1)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb2)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb3)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb1)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb2)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb3)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb1)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb2)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb3)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &max_bias)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &m0)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &m1)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_head_log2_val)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float), &logit_softcap)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &n_head_kv)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &mask_buffer)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &enable_mask)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb1_padded)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb2_padded)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb3_padded)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &mask_ne2)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &mask_ne3)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &mem_tex_sinks_1dbuf)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &enable_sinks)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &img_q_wmm)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &img_k_wmm)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem), &img_v_wmm)); + CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int), &d_head_q)); + + size_t global_work_size[3], local_work_size[3]; + + const int n_waves_v = d_head_q / 64; + + local_work_size[0] = 64; + local_work_size[1] = n_waves_v; + local_work_size[2] = 1; + + global_work_size[0] = 64; + global_work_size[1] = ((n_q + 64 - 1) / 64) * n_waves_v; + global_work_size[2] = n_batch * n_head; + + backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst); + + CL_CHECK(clReleaseMemObject(mem_tex_matrixO_1dbuf)); + CL_CHECK(clReleaseMemObject(img_q_wmm)); + CL_CHECK(clReleaseMemObject(img_k_wmm)); + CL_CHECK(clReleaseMemObject(img_v_wmm)); + + if (mem_tex_matrixMask_1dbuf) { + CL_CHECK(clReleaseMemObject(mem_tex_matrixMask_1dbuf)); + } + if (mem_matrixMask) { + CL_CHECK(clReleaseMemObject(mem_matrixMask)); + } + if (mem_matrixMask_padded) { + CL_CHECK(clReleaseMemObject(mem_matrixMask_padded)); + } + if (mem_tex_sinks_1dbuf) { + CL_CHECK(clReleaseMemObject(mem_tex_sinks_1dbuf)); + } + if (mem_sinksBuf) { + CL_CHECK(clReleaseMemObject(mem_sinksBuf)); + } + CL_CHECK(clReleaseMemObject(mem_matrixQ)); + CL_CHECK(clReleaseMemObject(mem_matrixK)); + CL_CHECK(clReleaseMemObject(mem_matrixV)); + CL_CHECK(clReleaseMemObject(mem_matrixO)); +} +#endif // GGML_OPENCL_USE_ADRENO_KERNELS + static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, const ggml_tensor * k, ggml_tensor * dst) { const ggml_tensor * v = dst->src[2]; const ggml_tensor * mask = dst->src[3]; @@ -17253,6 +17728,14 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0; const bool is_q4_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q4_0 && v->type == GGML_TYPE_Q4_0; +#ifdef GGML_OPENCL_USE_ADRENO_KERNELS + if (use_fa_bin_kernels_prefill(backend_ctx, q, k, v)) { + // We support the prefill path of flash attn with a specialized d_head = 64/128/256 + ggml_cl_flash_attn_prefill_bin(backend, q, k, dst); + return; + } +#endif + if (is_f16) { ggml_opencl_ensure_fa_variant(backend_ctx, d_head_q, d_head_v, FA_VARIANT_F16); } else if (is_mixed) { diff --git a/ggml/src/ggml-opencl/kernels/flash_attn_repack.cl b/ggml/src/ggml-opencl/kernels/flash_attn_repack.cl new file mode 100644 index 000000000..db78d5634 --- /dev/null +++ b/ggml/src/ggml-opencl/kernels/flash_attn_repack.cl @@ -0,0 +1,92 @@ +#pragma OPENCL EXTENSION cl_khr_fp16 : enable + +__kernel void kernel_repack_mask_for_wmm( + const global half* mask_buf, + const ulong mask_nb1, + const ulong mask_nb2, + const ulong mask_nb3, + const int mask_ne2, + global half* mask_buf_padded, + const ulong mask_nb1_padded, + const ulong mask_nb2_padded, + const ulong mask_nb3_padded +) { + int col = get_global_id(0); // 0 .. n_kv + int row = get_global_id(1); // 0 .. n_q + int slice = get_global_id(2); // 0 .. (n_head * n_batch) + + int head_idx = slice % mask_ne2; + int batch_idx = slice / mask_ne2; + + ulong src_off = (ulong)batch_idx * mask_nb3 + (ulong)head_idx * mask_nb2 + (ulong)row * mask_nb1; + ulong dst_off = (ulong)batch_idx * mask_nb3_padded + (ulong)head_idx * mask_nb2_padded + (ulong)row * mask_nb1_padded; + + mask_buf_padded[dst_off / 2 + col] = mask_buf[src_off / 2 + col]; +} + +__kernel void kernel_repack_q_for_wmm( + const global float* q_buf, + const ulong q_nb1, + const ulong q_nb2, + const ulong q_nb3, + const int n_head, + __write_only image3d_t img_q_wmm +) { + int k4 = get_global_id(0); + int row = get_global_id(1); + int slice = get_global_id(2); + int batch_idx = slice / n_head; + int head_idx = slice % n_head; + + + ulong elem_off = (batch_idx * q_nb3 + head_idx * q_nb2 + row * q_nb1) / 4 + (ulong)k4 * 4; + float4 v = vload4(elem_off / 4, q_buf); + + write_imageh(img_q_wmm, (int4)(row, slice, k4, 0), convert_half4(v)); +} + +__kernel void kernel_repack_k_for_wmm( + const global half* k_buf, + const ulong k_nb1, + const ulong k_nb2, + const ulong k_nb3, + const int n_head_kv, + const int n_kv, + __write_only image3d_t img_k_wmm +) { + int kk = get_global_id(0); + int row4 = get_global_id(1); + int slice = get_global_id(2); + int batch_idx = slice / n_head_kv; + int head_kv_idx = slice % n_head_kv; + + ulong base = batch_idx * k_nb3 + head_kv_idx * k_nb2; + int row0 = row4 * 4; + half4 v; + v.x = (row0 + 0 < n_kv) ? k_buf[(base + (ulong)(row0 + 0) * k_nb1) / 2 + kk] : (half)0; + v.y = (row0 + 1 < n_kv) ? k_buf[(base + (ulong)(row0 + 1) * k_nb1) / 2 + kk] : (half)0; + v.z = (row0 + 2 < n_kv) ? k_buf[(base + (ulong)(row0 + 2) * k_nb1) / 2 + kk] : (half)0; + v.w = (row0 + 3 < n_kv) ? k_buf[(base + (ulong)(row0 + 3) * k_nb1) / 2 + kk] : (half)0; + + write_imageh(img_k_wmm, (int4)(kk, row4, slice, 0), v); +} + +__kernel void kernel_repack_v_for_wmm( + const global half* v_buf, + const ulong v_nb1, + const ulong v_nb2, + const ulong v_nb3, + const int n_head_kv, + __write_only image3d_t img_v_wmm +) { + int hdim4 = get_global_id(0); // now fastest — walks contiguous memory + int row = get_global_id(1); + int slice = get_global_id(2); + int batch_idx = slice / n_head_kv; + int head_kv_idx = slice % n_head_kv; + + ulong row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + (ulong)row * v_nb1; + half4 v = vload4((row_off / 2 + (ulong)hdim4 * 4) / 4, v_buf); + + write_imageh(img_v_wmm, (int4)(row, hdim4, slice, 0), v); +} From b23701f77d47dad9de834d59ebfcbe25c9e8b46f Mon Sep 17 00:00:00 2001 From: TheArchitectit Date: Sat, 19 Sep 2026 00:32:52 -0500 Subject: [PATCH 6/9] cuda : fix CUB argsort corruption caused by in-place keys (#28389) argsort_f32_i32_cuda_cub called the one-shot DeviceRadixSort::SortPairs API with d_keys_in == d_keys_out (temp_keys, temp_keys). CUB's internal double-buffer ping-pong requires distinct key buffers: with aliased buffers the sort partially overwrites its own input mid-pass and emits a corrupted permutation, surfacing as intermittent garbage indices (e.g. backend top_k over a 248k-column vocab on Maxwell/CUDA 12.5/CCCL 2.x, which then triggered out-of-bounds gathers in downstream get_rows). Use a distinct keys-out buffer for all six call sites (plain and segmented, ascending and descending, size-query and execute). --------- Co-authored-by: Claude Opus 4.6 Co-authored-by: Oliver Simons --- ggml/src/ggml-cuda/argsort.cu | 27 +++++++++++++++------------ 1 file changed, 15 insertions(+), 12 deletions(-) diff --git a/ggml/src/ggml-cuda/argsort.cu b/ggml/src/ggml-cuda/argsort.cu index 26af90025..24115da09 100644 --- a/ggml/src/ggml-cuda/argsort.cu +++ b/ggml/src/ggml-cuda/argsort.cu @@ -51,9 +51,12 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool, cudaStream_t stream) { ggml_cuda_pool_alloc temp_indices_alloc(pool, ncols * nrows); ggml_cuda_pool_alloc temp_keys_alloc(pool, ncols * nrows); + // Device*Sort algorithms currently do not allow for in-place sorting/aliasing of input/outputs + ggml_cuda_pool_alloc temp_keys_out_alloc(pool, ncols * nrows); int * temp_indices = temp_indices_alloc.get(); float * temp_keys = temp_keys_alloc.get(); + float * temp_keys_out = temp_keys_out_alloc.get(); static const int block_size = 256; const dim3 grid_size((ncols + block_size - 1) / block_size, nrows); @@ -85,18 +88,18 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool, if (order == GGML_SORT_ORDER_ASC) { if (nrows == 1) { - CUDA_CHECK(DeviceRadixSort::SortPairs(nullptr, temp_storage_bytes, temp_keys, temp_keys, // keys (in-place) + CUDA_CHECK(DeviceRadixSort::SortPairs(nullptr, temp_storage_bytes, temp_keys, temp_keys_out, // keys in, keys out temp_indices, dst, // values (indices) ncols, 0, sizeof(float) * 8, stream)); } else if (is_capturing) { CUDA_CHECK(DeviceSegmentedRadixSort::SortPairs( - nullptr, temp_storage_bytes, temp_keys, temp_keys, // keys (in-place) + nullptr, temp_storage_bytes, temp_keys, temp_keys_out, // keys in, keys out temp_indices, dst, // values (indices) ncols * nrows, nrows, // num items, num segments offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream)); } else { CUDA_CHECK(DeviceSegmentedSort::SortPairs(nullptr, temp_storage_bytes, temp_keys, - temp_keys, // keys (in-place) + temp_keys_out, // keys out temp_indices, dst, // values (indices) ncols * nrows, nrows, // num items, num segments offset_iterator, offset_iterator + 1, stream)); @@ -104,15 +107,15 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool, } else { if (nrows == 1) { CUDA_CHECK(DeviceRadixSort::SortPairsDescending(nullptr, temp_storage_bytes, temp_keys, - temp_keys, // keys (in-place) + temp_keys_out, // keys out temp_indices, dst, // values (indices) ncols, 0, sizeof(float) * 8, stream)); } else if (is_capturing) { CUDA_CHECK(DeviceSegmentedRadixSort::SortPairsDescending( - nullptr, temp_storage_bytes, temp_keys, temp_keys, temp_indices, dst, ncols * nrows, nrows, + nullptr, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows, offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream)); } else { - CUDA_CHECK(DeviceSegmentedSort::SortPairsDescending(nullptr, temp_storage_bytes, temp_keys, temp_keys, + CUDA_CHECK(DeviceSegmentedSort::SortPairsDescending(nullptr, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows, offset_iterator, offset_iterator + 1, stream)); } @@ -124,31 +127,31 @@ void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool, if (order == GGML_SORT_ORDER_ASC) { if (nrows == 1) { CUDA_CHECK(DeviceRadixSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, - temp_keys, // keys (in-place) + temp_keys_out, // keys out temp_indices, dst, // values (indices) ncols, 0, sizeof(float) * 8, stream)); } else if (is_capturing) { - CUDA_CHECK(DeviceSegmentedRadixSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys, + CUDA_CHECK(DeviceSegmentedRadixSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows, offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream)); } else { - CUDA_CHECK(DeviceSegmentedSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys, + CUDA_CHECK(DeviceSegmentedSort::SortPairs(d_temp_storage, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows, offset_iterator, offset_iterator + 1, stream)); } } else { if (nrows == 1) { CUDA_CHECK(DeviceRadixSort::SortPairsDescending(d_temp_storage, temp_storage_bytes, temp_keys, - temp_keys, // keys (in-place) + temp_keys_out, // keys out temp_indices, dst, // values (indices) ncols, 0, sizeof(float) * 8, stream)); } else if (is_capturing) { CUDA_CHECK(DeviceSegmentedRadixSort::SortPairsDescending( - d_temp_storage, temp_storage_bytes, temp_keys, temp_keys, temp_indices, dst, ncols * nrows, nrows, + d_temp_storage, temp_storage_bytes, temp_keys, temp_keys_out, temp_indices, dst, ncols * nrows, nrows, offset_iterator, offset_iterator + 1, 0, sizeof(float) * 8, stream)); } else { CUDA_CHECK(DeviceSegmentedSort::SortPairsDescending(d_temp_storage, temp_storage_bytes, temp_keys, - temp_keys, temp_indices, dst, ncols * nrows, nrows, + temp_keys_out, temp_indices, dst, ncols * nrows, nrows, offset_iterator, offset_iterator + 1, stream)); } } From 59fc5a1ca3842241dd53617ae2ae030c1a015061 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Sat, 19 Sep 2026 11:27:30 +0300 Subject: [PATCH 7/9] metal : support qwen4exp hc ops (#29000) Add support for the new DSV4 HC op variants used by qwen4exp: - hc_pre with per-element sigmoid gate (gated variant) - hc_post with identity mixing (comb == nullptr) Assisted-by: pi:llama.cpp/Qwen3.8-27B --- ggml/src/ggml-metal/ggml-metal-device.cpp | 27 +++++++-- ggml/src/ggml-metal/ggml-metal-device.h | 2 +- ggml/src/ggml-metal/ggml-metal-device.m | 10 +--- ggml/src/ggml-metal/ggml-metal-impl.h | 2 + ggml/src/ggml-metal/ggml-metal-ops.cpp | 19 ++++--- ggml/src/ggml-metal/kernels/misc.metal | 68 ++++++++++++++++++++++- 6 files changed, 106 insertions(+), 22 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index b510cb957..0dcfad3af 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -496,14 +496,29 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexe return res; } -ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_metal_library_t lib, ggml_op op) { +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc(ggml_metal_library_t lib, const ggml_tensor * op) { const char * name = nullptr; - switch (op) { - case GGML_OP_DSV4_HC_COMB: name = "kernel_dsv4_hc_comb_f32"; break; - case GGML_OP_DSV4_HC_PRE: name = "kernel_dsv4_hc_pre_f32"; break; - case GGML_OP_DSV4_HC_POST: name = "kernel_dsv4_hc_post_f32"; break; - default: GGML_ABORT("fatal error"); + switch (op->op) { + case GGML_OP_DSV4_HC_COMB: + name = "kernel_dsv4_hc_comb_f32"; + break; + case GGML_OP_DSV4_HC_PRE: + if (ggml_get_op_params_i32(op, 1) != 0) { + name = "kernel_dsv4_hc_pre_gated_f32"; + } else { + name = "kernel_dsv4_hc_pre_f32"; + } + break; + case GGML_OP_DSV4_HC_POST: + if (op->src[3]) { + name = "kernel_dsv4_hc_post_f32"; + } else { + name = "kernel_dsv4_hc_post_nocomb_f32"; + } + break; + default: + GGML_ABORT("fatal error"); } ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index f6243ffbd..0514f9ef0 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -126,7 +126,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_ad struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, const struct ggml_tensor * op); -struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, enum ggml_op op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op, bool tail); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index c734c8e13..952d1c0a6 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1802,8 +1802,6 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te op->src[1]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32 && op->src[0]->ne[1] == 4 && - op->src[1]->ne[0] == 4 && - op->src[1]->ne[2] == 1 && ggml_is_contiguous_rows(op->src[0]) && ggml_is_contiguous_rows(op->src[1]); case GGML_OP_DSV4_HC_POST: @@ -1811,17 +1809,15 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && op->src[2]->type == GGML_TYPE_F32 && - op->src[3] != NULL && - op->src[3]->type == GGML_TYPE_F32 && + (op->src[3] == NULL || op->src[3]->type == GGML_TYPE_F32) && op->type == GGML_TYPE_F32 && op->src[1]->ne[1] == 4 && op->src[2]->ne[0] == 4 && - op->src[3]->ne[0] == 4 && - op->src[3]->ne[1] == 4 && + (op->src[3] == NULL || (op->src[3]->ne[0] == 4 && op->src[3]->ne[1] == 4)) && ggml_is_contiguous_rows(op->src[0]) && ggml_is_contiguous_rows(op->src[1]) && ggml_is_contiguous_rows(op->src[2]) && - ggml_is_contiguous_rows(op->src[3]); + (op->src[3] == NULL || ggml_is_contiguous_rows(op->src[3])); case GGML_OP_SSM_SCAN: return has_simdgroup_reduction; case GGML_OP_SSM_CONV: diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index 7a2c65aaa..d84ca937b 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -1283,8 +1283,10 @@ typedef struct { uint64_t nb_x2; uint64_t nb_w0; uint64_t nb_w1; + uint64_t nb_w2; uint64_t nb_d0; uint64_t nb_d1; + float scale; } ggml_metal_kargs_dsv4_hc_pre; typedef struct { diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index cc1bebfaa..77c399bdb 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1405,7 +1405,7 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); ggml_metal_encoder_t enc = ctx->enc; - auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc(ctx->lib, op->op); + auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc(ctx->lib, op); ggml_metal_encoder_set_pipeline(enc, pipeline); @@ -1467,8 +1467,10 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { /*.nb_x2 =*/ x->nb[2], /*.nb_w0 =*/ weights->nb[0], /*.nb_w1 =*/ weights->nb[1], + /*.nb_w2 =*/ weights->nb[2], /*.nb_d0 =*/ op->nb[0], /*.nb_d1 =*/ op->nb[1], + /*.scale =*/ ggml_get_op_params_f32(op, 0), }; ggml_metal_encoder_set_bytes (enc, &args, sizeof(args), 0); @@ -1491,7 +1493,6 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { GGML_ASSERT(x->type == GGML_TYPE_F32); GGML_ASSERT(residual->type == GGML_TYPE_F32); GGML_ASSERT(post->type == GGML_TYPE_F32); - GGML_ASSERT(comb->type == GGML_TYPE_F32); GGML_ASSERT(op->type == GGML_TYPE_F32); GGML_ASSERT(residual->ne[1] == 4); @@ -1505,9 +1506,9 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { /*.nb_r2 =*/ residual->nb[2], /*.nb_p0 =*/ post->nb[0], /*.nb_p1 =*/ post->nb[1], - /*.nb_c0 =*/ comb->nb[0], - /*.nb_c1 =*/ comb->nb[1], - /*.nb_c2 =*/ comb->nb[2], + /*.nb_c0 =*/ comb ? comb->nb[0] : 0, + /*.nb_c1 =*/ comb ? comb->nb[1] : 0, + /*.nb_c2 =*/ comb ? comb->nb[2] : 0, /*.nb_d0 =*/ op->nb[0], /*.nb_d1 =*/ op->nb[1], /*.nb_d2 =*/ op->nb[2], @@ -1517,8 +1518,12 @@ int ggml_metal_op_dsv4_hc(ggml_metal_op_t ctx, int idx) { ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(x), 1); ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(residual), 2); ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(post), 3); - ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(comb), 4); - ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); + if (comb) { + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(comb), 4); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); + } else { + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 4); + } const int n_tiles = (args.n_embd + 31)/32; const int nsg = std::min(4, n_tiles); diff --git a/ggml/src/ggml-metal/kernels/misc.metal b/ggml/src/ggml-metal/kernels/misc.metal index 11104b4d8..15a18e04a 100644 --- a/ggml/src/ggml-metal/kernels/misc.metal +++ b/ggml/src/ggml-metal/kernels/misc.metal @@ -531,7 +531,73 @@ kernel void kernel_dsv4_hc_pre_f32( result = fma(*(device const float *) (xb + ih*args.nb_x1), w[ih], result); } - *(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = result; + *(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = args.scale*result; +} + +kernel void kernel_dsv4_hc_pre_gated_f32( + constant ggml_metal_kargs_dsv4_hc_pre & args, + device const char * x, + device const char * gate, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr ushort hc = 4; + + const int it = tgpig.y; + const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg; + + if (i0 >= args.n_embd) { + return; + } + + device const char * xb = x + i0*args.nb_x0 + it*args.nb_x2; + device const char * gb = gate + i0*args.nb_w0 + it*args.nb_w2; + float result = 0.0f; + FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) { + const float g = 1.0f/(1.0f + exp(-*(device const float *) (gb + ih*args.nb_w1))); + result = fma(*(device const float *) (xb + ih*args.nb_x1), g, result); + } + + *(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = args.scale*result; +} + +kernel void kernel_dsv4_hc_post_nocomb_f32( + constant ggml_metal_kargs_dsv4_hc_post & args, + device const char * x, + device const char * residual, + device const char * post, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr ushort hc = 4; + + const int it = tgpig.y; + const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg; + + float post_lane = 0.0f; + if (tiisg < hc) { + post_lane = *(device const float *) (post + tiisg*args.nb_p0 + it*args.nb_p1); + } + + float post_reg[hc]; + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + post_reg[idst] = simd_shuffle(post_lane, idst); + } + + if (i0 >= args.n_embd) { + return; + } + + const float xv = *(device const float *) (x + i0*args.nb_x0 + it*args.nb_x1); + device const char * rb = residual + i0*args.nb_r0 + it*args.nb_r2; + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + const float rv = *(device const float *) (rb + idst*args.nb_r1); + *(device float *) (dst + i0*args.nb_d0 + idst*args.nb_d1 + it*args.nb_d2) = xv*post_reg[idst] + rv; + } } kernel void kernel_dsv4_hc_post_f32( From efa28e950ea3a41648aaf354b3a743dd4708f954 Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Sat, 19 Sep 2026 11:27:46 +0300 Subject: [PATCH 8/9] test-llama-archs : generate dummy test vocab (#29084) Assisted-by: pi:llama.cpp/DeepSeek-V4-Flash-Vision-Exp --- include/llama.h | 1 + src/llama-model-saver.cpp | 26 ++++++++--------- src/llama-vocab.cpp | 58 ++++++++++++++++++++++++++++++++++++++ tests/test-llama-archs.cpp | 14 ++++++++- 4 files changed, 85 insertions(+), 14 deletions(-) diff --git a/include/llama.h b/include/llama.h index ac2215dc7..31bbf8b0d 100644 --- a/include/llama.h +++ b/include/llama.h @@ -77,6 +77,7 @@ extern "C" { LLAMA_VOCAB_TYPE_UGM = 4, // T5 tokenizer based on Unigram LLAMA_VOCAB_TYPE_RWKV = 5, // RWKV tokenizer based on greedy tokenization LLAMA_VOCAB_TYPE_PLAMO2 = 6, // PLaMo-2 tokenizer based on Aho-Corasick with dynamic programming + LLAMA_VOCAB_TYPE_TEST = 7, // Dummy tokenizer for testing: rolling hash of fixed-size chunks -> tokens, tokens -> hex }; enum llama_rope_type { diff --git a/src/llama-model-saver.cpp b/src/llama-model-saver.cpp index 0f5155b2e..39160a417 100644 --- a/src/llama-model-saver.cpp +++ b/src/llama-model-saver.cpp @@ -387,13 +387,13 @@ void llama_model_saver::add_kv_from_model() { add_kv(LLM_KV_TOKENIZER_SCORES, scores); add_kv(LLM_KV_TOKENIZER_MERGES, vocab.get_bpe_merges()); // FIXME llama_token is type i32 but when reading in a GGUF file u32 is expected, not an issue for writing though - add_kv(LLM_KV_TOKENIZER_BOS_ID, uint32_t(vocab.token_bos())); - add_kv(LLM_KV_TOKENIZER_EOS_ID, uint32_t(vocab.token_eos())); - add_kv(LLM_KV_TOKENIZER_EOT_ID, uint32_t(vocab.token_eot())); - add_kv(LLM_KV_TOKENIZER_EOM_ID, uint32_t(vocab.token_eom())); - add_kv(LLM_KV_TOKENIZER_UNK_ID, uint32_t(vocab.token_unk())); - add_kv(LLM_KV_TOKENIZER_SEP_ID, uint32_t(vocab.token_sep())); - add_kv(LLM_KV_TOKENIZER_PAD_ID, uint32_t(vocab.token_pad())); + if (vocab.token_bos() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_BOS_ID, uint32_t(vocab.token_bos())); } + if (vocab.token_eos() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_EOS_ID, uint32_t(vocab.token_eos())); } + if (vocab.token_eot() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_EOT_ID, uint32_t(vocab.token_eot())); } + if (vocab.token_eom() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_EOM_ID, uint32_t(vocab.token_eom())); } + if (vocab.token_unk() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_UNK_ID, uint32_t(vocab.token_unk())); } + if (vocab.token_sep() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_SEP_ID, uint32_t(vocab.token_sep())); } + if (vocab.token_pad() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_PAD_ID, uint32_t(vocab.token_pad())); } // add_kv(LLM_KV_TOKENIZER_CLS_ID, uint32_t(vocab.token_bos())); // deprecated // add_kv(LLM_KV_TOKENIZER_MASK_ID, ???); add_kv(LLM_KV_TOKENIZER_ADD_BOS, vocab.get_add_bos()); @@ -404,12 +404,12 @@ void llama_model_saver::add_kv_from_model() { add_kv(LLM_KV_TOKENIZER_PRECOMPILED_CHARSMAP, vocab.get_precompiled_charsmap()); // add_kv(LLM_KV_TOKENIZER_HF_JSON, ???); // add_kv(LLM_KV_TOKENIZER_RWKV, ???); - add_kv(LLM_KV_TOKENIZER_FIM_PRE_ID, uint32_t(vocab.token_fim_pre())); - add_kv(LLM_KV_TOKENIZER_FIM_SUF_ID, uint32_t(vocab.token_fim_suf())); - add_kv(LLM_KV_TOKENIZER_FIM_MID_ID, uint32_t(vocab.token_fim_mid())); - add_kv(LLM_KV_TOKENIZER_FIM_PAD_ID, uint32_t(vocab.token_fim_pad())); - add_kv(LLM_KV_TOKENIZER_FIM_REP_ID, uint32_t(vocab.token_fim_rep())); - add_kv(LLM_KV_TOKENIZER_FIM_SEP_ID, uint32_t(vocab.token_fim_sep())); + if (vocab.token_fim_pre() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_FIM_PRE_ID, uint32_t(vocab.token_fim_pre())); } + if (vocab.token_fim_suf() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_FIM_SUF_ID, uint32_t(vocab.token_fim_suf())); } + if (vocab.token_fim_mid() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_FIM_MID_ID, uint32_t(vocab.token_fim_mid())); } + if (vocab.token_fim_pad() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_FIM_PAD_ID, uint32_t(vocab.token_fim_pad())); } + if (vocab.token_fim_rep() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_FIM_REP_ID, uint32_t(vocab.token_fim_rep())); } + if (vocab.token_fim_sep() != LLAMA_TOKEN_NULL) { add_kv(LLM_KV_TOKENIZER_FIM_SEP_ID, uint32_t(vocab.token_fim_sep())); } // TODO: implement LoRA support // add_kv(LLM_KV_ADAPTER_TYPE, ???); diff --git a/src/llama-vocab.cpp b/src/llama-vocab.cpp index 737e07275..e038637ce 100644 --- a/src/llama-vocab.cpp +++ b/src/llama-vocab.cpp @@ -2087,6 +2087,16 @@ void llama_vocab::impl::load(llama_model_loader & ml, const LLM_KV & kv) { special_unk_id = LLAMA_TOKEN_NULL; special_sep_id = LLAMA_TOKEN_NULL; special_pad_id = LLAMA_TOKEN_NULL; + } else if (tokenizer_model == "test") { + type = LLAMA_VOCAB_TYPE_TEST; + + // default special tokens + special_bos_id = LLAMA_TOKEN_NULL; + special_eos_id = LLAMA_TOKEN_NULL; + special_unk_id = LLAMA_TOKEN_NULL; + special_sep_id = LLAMA_TOKEN_NULL; + special_pad_id = LLAMA_TOKEN_NULL; + special_mask_id = LLAMA_TOKEN_NULL; } else if (tokenizer_model == "plamo2") { type = LLAMA_VOCAB_TYPE_PLAMO2; @@ -3134,6 +3144,7 @@ std::string llama_vocab::impl::type_name() const{ case LLAMA_VOCAB_TYPE_UGM: return "UGM"; case LLAMA_VOCAB_TYPE_RWKV: return "RWKV"; case LLAMA_VOCAB_TYPE_PLAMO2: return "PLaMo2"; + case LLAMA_VOCAB_TYPE_TEST: return "TEST"; default: return "unknown"; } } @@ -3222,6 +3233,9 @@ void llama_vocab::impl::init_tokenizer(enum llama_vocab_type type) { case LLAMA_VOCAB_TYPE_PLAMO2: tokenizer = std::make_unique(vocab); break; + case LLAMA_VOCAB_TYPE_TEST: + tokenizer = std::make_unique(); + break; default: GGML_ABORT("unsupported vocab type"); } @@ -3595,6 +3609,42 @@ std::vector llama_vocab::impl::tokenize( } } } break; + case LLAMA_VOCAB_TYPE_TEST: + { + const uint32_t n_vocab = vocab.n_tokens(); + constexpr size_t chunk_size = 5; + + // reserve output to avoid repeated reallocations + size_t n_tokens = 0; + for (const auto & fragment : fragment_buffer) { + if (fragment.type == FRAGMENT_BUFFER_VARIANT_TYPE_RAW_TEXT) { + n_tokens += (fragment.length + chunk_size - 1) / chunk_size; + } else { + ++n_tokens; + } + } + output.reserve(output.size() + n_tokens); + + for (const auto & fragment : fragment_buffer) { + if (fragment.type == FRAGMENT_BUFFER_VARIANT_TYPE_RAW_TEXT) { + const auto & text = fragment.raw_text; + const size_t begin = fragment.offset; + const size_t end = begin + fragment.length; + size_t pos = begin; + while (pos < end) { + const size_t n = std::min(chunk_size, end - pos); + uint64_t hash = 0; + for (size_t i = 0; i < n; ++i) { + hash = hash*31 + (uint8_t) text[pos + i]; + } + output.push_back((llama_token)(hash % n_vocab)); + pos += n; + } + } else { // if (fragment.type == FRAGMENT_BUFFER_VARIANT_TYPE_TOKEN) + output.push_back(fragment.token); + } + } + } break; case LLAMA_VOCAB_TYPE_NONE: GGML_ABORT("fatal error"); } @@ -3693,6 +3743,11 @@ int32_t llama_vocab::impl::token_to_piece(llama_token token, char * buf, int32_t memcpy(buf, result.data(), result.size()); return (int)result.size(); } + case LLAMA_VOCAB_TYPE_TEST: { + // tokens -> text: simply stringify the token id in hex + std::string result = format("%x", token); + return _try_copy(result.data(), result.size()); + } case LLAMA_VOCAB_TYPE_PLAMO2: { // PLaMo-2 uses similar token handling as BPE/SPM if (vocab.is_byte(token)) { @@ -3963,6 +4018,9 @@ llama_token llama_vocab::byte_to_token(uint8_t ch) const { snprintf(hex_str, sizeof(hex_str), "<0x%02X>", ch); return pimpl->token_to_id.at(hex_str); } + case LLAMA_VOCAB_TYPE_TEST: + // TEST tokens have no byte-level mapping + return LLAMA_TOKEN_NULL; default: GGML_ABORT("fatal error"); } diff --git a/tests/test-llama-archs.cpp b/tests/test-llama-archs.cpp index 568f7234c..f848fc139 100644 --- a/tests/test-llama-archs.cpp +++ b/tests/test-llama-archs.cpp @@ -338,7 +338,19 @@ static gguf_context_ptr get_gguf_ctx(const llm_arch arch, const bool moe) { ms.add_kv(LLM_KV_SWIGLU_CLAMP_EXP, 7.0f); } - ms.add_kv(LLM_KV_TOKENIZER_MODEL, "no_vocab"); + // dummy tokenizer: token ids are derived from fixed-size chunks and detokenized as hex ids + { + std::vector tokenizer_list(n_vocab); + std::vector tokenizer_scores(n_vocab, 0.0f); + + ms.add_kv(LLM_KV_TOKENIZER_MODEL, "test"); + for (uint32_t i = 0; i < n_vocab; i++) { + tokenizer_list[i] = "tok_" + std::to_string(i); + } + ms.add_kv(LLM_KV_TOKENIZER_LIST, tokenizer_list); + ms.add_kv(LLM_KV_TOKENIZER_SCORES, tokenizer_scores); + } + // ms.add_kv(LLM_KV_DENSE_2_FEAT_OUT, n_embd); // ms.add_kv(LLM_KV_DENSE_3_FEAT_IN, n_embd); From 60b06ab9a9eeec26f8125c9316ccbf4ee4713d1f Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Sat, 19 Sep 2026 11:33:03 +0300 Subject: [PATCH 9/9] metal : fix FA support checks (#29122) --- ggml/src/ggml-metal/ggml-metal-device.m | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 952d1c0a6..0f42d5700 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1733,6 +1733,12 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te op->src[0]->ne[0] != 576) { return false; } + if (op->src[1]->ne[0] == 72 && op->src[1]->ne[0] != op->src[2]->ne[0]) { + return false; + } + if (op->src[1]->ne[0] < op->src[2]->ne[0]) { + return false; + } if (op->src[1]->type != op->src[2]->type) { return false; }