From 7d701b59296bc6f5ba504d5a4cddaf416c449a15 Mon Sep 17 00:00:00 2001 From: lhez Date: Mon, 7 Sep 2026 23:26:34 -0700 Subject: [PATCH] opencl: properly handle non-contiguous inputs to conv2d (#28503) * opencl: fix conv2d non-contiguous strides * opencl: format --- ggml/src/ggml-opencl/ggml-opencl.cpp | 77 ++++++++++++++----- ggml/src/ggml-opencl/kernels/conv2d.cl | 8 +- .../src/ggml-opencl/kernels/conv2d_f16_f32.cl | 8 +- 3 files changed, 66 insertions(+), 27 deletions(-) diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp index d737aea12..3002835e8 100644 --- a/ggml/src/ggml-opencl/ggml-opencl.cpp +++ b/ggml/src/ggml-opencl/ggml-opencl.cpp @@ -17906,16 +17906,34 @@ static void ggml_cl_conv_2d(ggml_backend_t backend, const ggml_tensor * src0, co cl_ulong offset1 = extra1->offset + src1->view_offs; cl_ulong offsetd = extrad->offset + dst->view_offs; - const cl_uint Cout = ne03; const cl_uint Cin = ne02; const cl_uint N = ne13; - const cl_uint KW = ne00; const cl_uint KH = ne01; const cl_uint W = ne10; const cl_uint H = ne11; const cl_uint OW = ne0; const cl_uint OH = ne1; + const cl_uint Cout = ne03; + const cl_uint Cin = ne02; + const cl_uint N = ne13; + const cl_uint KW = ne00; + const cl_uint KH = ne01; + const cl_uint W = ne10; + const cl_uint H = ne11; + const cl_uint OW = ne0; + const cl_uint OH = ne1; - const cl_uint s0 = dst->op_params[0]; const cl_uint s1 = dst->op_params[1]; - const cl_uint p0 = dst->op_params[2]; const cl_uint p1 = dst->op_params[3]; - const cl_uint d0 = dst->op_params[4]; const cl_uint d1 = dst->op_params[5]; + const cl_uint s0 = dst->op_params[0]; + const cl_uint s1 = dst->op_params[1]; + const cl_uint p0 = dst->op_params[2]; + const cl_uint p1 = dst->op_params[3]; + const cl_uint d0 = dst->op_params[4]; + const cl_uint d1 = dst->op_params[5]; - const cl_uint cl_nb01 = nb01/ggml_type_size(src0->type); const cl_uint cl_nb02 = nb02/ggml_type_size(src0->type); const cl_uint cl_nb03 = nb03/ggml_type_size(src0->type); - const cl_uint cl_nb11 = nb11/ggml_type_size(src1->type); const cl_uint cl_nb12 = nb12/ggml_type_size(src1->type); const cl_uint cl_nb13 = nb13/ggml_type_size(src1->type); - const cl_uint cl_nb1 = nb1/ggml_type_size(dst->type); const cl_uint cl_nb2 = nb2/ggml_type_size(dst->type); const cl_uint cl_nb3 = nb3/ggml_type_size(dst->type); + const cl_uint cl_nb00 = nb00/ggml_type_size(src0->type); + const cl_uint cl_nb01 = nb01/ggml_type_size(src0->type); + const cl_uint cl_nb02 = nb02/ggml_type_size(src0->type); + const cl_uint cl_nb03 = nb03/ggml_type_size(src0->type); + const cl_uint cl_nb10 = nb10/ggml_type_size(src1->type); + const cl_uint cl_nb11 = nb11/ggml_type_size(src1->type); + const cl_uint cl_nb12 = nb12/ggml_type_size(src1->type); + const cl_uint cl_nb13 = nb13/ggml_type_size(src1->type); + const cl_uint cl_nb1 = nb1/ggml_type_size(dst->type); + const cl_uint cl_nb2 = nb2/ggml_type_size(dst->type); + const cl_uint cl_nb3 = nb3/ggml_type_size(dst->type); const int64_t NPQ = (int64_t)N * OW * OH; @@ -17951,18 +17969,39 @@ static void ggml_cl_conv_2d(ggml_backend_t backend, const ggml_tensor * src0, co } cl_uint idx = 0; - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_mem), &extra0->data_device)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_ulong), &offset0)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_mem), &extra1->data_device)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_ulong), &offset1)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_mem), &extrad->data_device)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_ulong), &offsetd)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_mem), &extra0->data_device)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_ulong), &offset0)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_mem), &extra1->data_device)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_ulong), &offset1)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_mem), &extrad->data_device)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_ulong), &offsetd)); CL_CHECK(clSetKernelArg(kernel, idx++, shmem_size, NULL)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &Cout)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &Cin)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &N)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &KW)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &KH)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &W)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &H)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &OW)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &OH)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &s0)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &s1)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &p0)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &p1)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &d0)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &d1)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb01)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb02)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb03)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb11)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb12)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb13)); - CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb1)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb2)); CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb3)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &Cout)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &Cin)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &N)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &KW)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &KH)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &W)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &H)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &OW)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &OH)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &s0)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &s1)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &p0)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &p1)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &d0)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &d1)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb00)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb01)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb02)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb03)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb10)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb11)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb12)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb13)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb1)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb2)); + CL_CHECK(clSetKernelArg(kernel, idx++, sizeof(cl_uint), &cl_nb3)); size_t global_work_size[] = { (size_t)NB_K * WG_K, (size_t)NB_NPQ * WG_NPQ, 1 }; size_t local_work_size[] = { (size_t)WG_K, (size_t)WG_NPQ, 1 }; diff --git a/ggml/src/ggml-opencl/kernels/conv2d.cl b/ggml/src/ggml-opencl/kernels/conv2d.cl index e339c90cf..8a04c2e59 100644 --- a/ggml/src/ggml-opencl/kernels/conv2d.cl +++ b/ggml/src/ggml-opencl/kernels/conv2d.cl @@ -48,8 +48,8 @@ kernel void kernel_conv_2d( uint Cout, uint Cin, uint N, uint KW, uint KH, uint W, uint H, uint OW, uint OH, uint s0, uint s1, uint p0, uint p1, uint d0, uint d1, - uint nb01, uint nb02, uint nb03, - uint nb11, uint nb12, uint nb13, + uint nb00, uint nb01, uint nb02, uint nb03, + uint nb10, uint nb11, uint nb12, uint nb13, uint nb1, uint nb2, uint nb3 ) { global T_FLOAT* knl_data = (global T_FLOAT*) ((global char*)p_knl + off_knl); @@ -95,7 +95,7 @@ kernel void kernel_conv_2d( const uint Cin_idx = crs_g / (KW*KH); const uint KH_idx = (crs_g - Cin_idx*KW*KH) / KW; const uint KW_idx = crs_g - Cin_idx*KW*KH - KH_idx*KW; - const uint knl_idx = KW_idx + KH_idx*nb01 + Cin_idx*nb02 + k_g*nb03; + const uint knl_idx = KW_idx*nb00 + KH_idx*nb01 + Cin_idx*nb02 + k_g*nb03; Ash[k_l * BS_CRS + crs_l] = knl_data[knl_idx]; } else { Ash[k_l * BS_CRS + crs_l] = (T_FLOAT)0.0f; @@ -123,7 +123,7 @@ kernel void kernel_conv_2d( const int W_idx = (int)(OW_idx * s0 + KW_idx * d0 - p0); if (H_idx >= 0 && H_idx < H && W_idx >= 0 && W_idx < W) { - const uint src_idx = W_idx + H_idx * nb11 + Cin_idx * nb12 + N_idx * nb13; + const uint src_idx = W_idx * nb10 + H_idx * nb11 + Cin_idx * nb12 + N_idx * nb13; ((T_FLOAT*)&val)[v] = src_data[src_idx]; } } diff --git a/ggml/src/ggml-opencl/kernels/conv2d_f16_f32.cl b/ggml/src/ggml-opencl/kernels/conv2d_f16_f32.cl index cb05637f3..94788e7e0 100644 --- a/ggml/src/ggml-opencl/kernels/conv2d_f16_f32.cl +++ b/ggml/src/ggml-opencl/kernels/conv2d_f16_f32.cl @@ -39,8 +39,8 @@ kernel void kernel_conv_2d( uint Cout, uint Cin, uint N, uint KW, uint KH, uint W, uint H, uint OW, uint OH, uint s0, uint s1, uint p0, uint p1, uint d0, uint d1, - uint nb01, uint nb02, uint nb03, - uint nb11, uint nb12, uint nb13, + uint nb00, uint nb01, uint nb02, uint nb03, + uint nb10, uint nb11, uint nb12, uint nb13, uint nb1, uint nb2, uint nb3 ) { global half* knl_data = (global half*) ((global char*)p_knl + off_knl); @@ -86,7 +86,7 @@ kernel void kernel_conv_2d( const uint Cin_idx = crs_g / (KW*KH); const uint KH_idx = (crs_g - Cin_idx*KW*KH) / KW; const uint KW_idx = crs_g - Cin_idx*KW*KH - KH_idx*KW; - const uint knl_idx = KW_idx + KH_idx*nb01 + Cin_idx*nb02 + k_g*nb03; + const uint knl_idx = KW_idx*nb00 + KH_idx*nb01 + Cin_idx*nb02 + k_g*nb03; Ash[k_l * BS_CRS + crs_l] = knl_data[knl_idx]; } else { Ash[k_l * BS_CRS + crs_l] = (half)0.0f; @@ -114,7 +114,7 @@ kernel void kernel_conv_2d( const int W_idx = (int)(OW_idx * s0 + KW_idx * d0 - p0); if (H_idx >= 0 && H_idx < H && W_idx >= 0 && W_idx < W) { - const uint src_idx = W_idx + H_idx * nb11 + Cin_idx * nb12 + N_idx * nb13; + const uint src_idx = W_idx * nb10 + H_idx * nb11 + Cin_idx * nb12 + N_idx * nb13; ((float*)&val)[v] = src_data[src_idx]; } }