Commit 7d701b592 for llama.cpp
commit 7d701b59296bc6f5ba504d5a4cddaf416c449a15
Author: lhez <lih@qti.qualcomm.com>
Date: Mon Sep 7 23:26:34 2026 -0700
opencl: properly handle non-contiguous inputs to conv2d (#28503)
* opencl: fix conv2d non-contiguous strides
* opencl: format
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 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 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 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];
}
}