opencl: transpose q4_K noshuffle scales for coalesced reads (#25805)

This commit is contained in:
Hongqiang Wang
2026-07-17 07:49:43 -07:00
committed by GitHub
parent 7d56da7e54
commit 86d86ed439
4 changed files with 49 additions and 38 deletions
+15 -6
View File
@@ -6012,7 +6012,8 @@ static void transpose_2d(
cl_kernel kernel, cl_kernel kernel,
cl_mem src, cl_mem dst, size_t size, cl_mem src, cl_mem dst, size_t size,
cl_int stride, cl_int rows, cl_int stride, cl_int rows,
bool blocking = true bool blocking = true,
bool auto_local = false // let driver pick local size for non-uniform workgroups
) { ) {
static ggml_cl_buffer buf; static ggml_cl_buffer buf;
@@ -6038,7 +6039,7 @@ static void transpose_2d(
size_t local_size[3] = {64, 1, 1}; size_t local_size[3] = {64, 1, 1};
size_t global_size[3] = {(size_t)stride, (size_t)rows, 1};; size_t global_size[3] = {(size_t)stride, (size_t)rows, 1};;
CL_CHECK(clEnqueueNDRangeKernel(backend_ctx->queue, kernel, 3, NULL, CL_CHECK(clEnqueueNDRangeKernel(backend_ctx->queue, kernel, 3, NULL,
global_size, local_size, 0, NULL, NULL)); global_size, auto_local ? NULL : local_size, 0, NULL, NULL));
if (blocking) { if (blocking) {
CL_CHECK(clEnqueueCopyBuffer(backend_ctx->queue, trans, dst, 0, 0, size, 0, NULL, &evt)); CL_CHECK(clEnqueueCopyBuffer(backend_ctx->queue, trans, dst, 0, 0, size, 0, NULL, &evt));
@@ -6055,10 +6056,11 @@ static void transpose_2d_as_8b(
ggml_backend_opencl_context * backend_ctx, ggml_backend_opencl_context * backend_ctx,
cl_mem src, cl_mem dst, size_t size, cl_mem src, cl_mem dst, size_t size,
cl_int stride, cl_int rows, cl_int stride, cl_int rows,
bool blocking = true bool blocking = true,
bool auto_local = false
) { ) {
transpose_2d(backend_ctx, backend_ctx->kernel_transpose_8_buf, transpose_2d(backend_ctx, backend_ctx->kernel_transpose_8_buf,
src, dst, size, stride, rows, blocking); src, dst, size, stride, rows, blocking, auto_local);
} }
static void transpose_2d_as_16b( static void transpose_2d_as_16b(
@@ -9054,6 +9056,9 @@ static void ggml_backend_opencl_buffer_set_tensor(ggml_backend_buffer_t buffer,
transpose_2d_as_16b(backend_ctx, extra->q, extra->q, size_q, K/4, M); transpose_2d_as_16b(backend_ctx, extra->q, extra->q, size_q, K/4, M);
transpose_2d_as_16b(backend_ctx, extra->d, extra->d, size_d, K/256, M); transpose_2d_as_16b(backend_ctx, extra->d, extra->d, size_d, K/256, M);
transpose_2d_as_16b(backend_ctx, extra->dm, extra->dm, size_dm, K/256, M); transpose_2d_as_16b(backend_ctx, extra->dm, extra->dm, size_dm, K/256, M);
// Transpose s as uchar
transpose_2d_as_8b(backend_ctx, extra->s, extra->s, size_s, K/256*12, M, true, true);
} }
#endif // GGML_OPENCL_USE_ADRENO_KERNELS #endif // GGML_OPENCL_USE_ADRENO_KERNELS
return; return;
@@ -10222,23 +10227,27 @@ static void ggml_backend_opencl_buffer_get_tensor(ggml_backend_buffer_t buffer,
size_t size_q = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*ggml_blck_size(tensor->type)/2; size_t size_q = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*ggml_blck_size(tensor->type)/2;
size_t size_d = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*sizeof(ggml_fp16_t); size_t size_d = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*sizeof(ggml_fp16_t);
size_t size_dm = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*sizeof(ggml_fp16_t); size_t size_dm = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*sizeof(ggml_fp16_t);
size_t size_s = ggml_nelements(tensor)/ggml_blck_size(tensor->type)*12;
static ggml_cl_buffer buf_trans_q; static ggml_cl_buffer buf_trans_q;
static ggml_cl_buffer buf_trans_d; static ggml_cl_buffer buf_trans_d;
static ggml_cl_buffer buf_trans_dm; static ggml_cl_buffer buf_trans_dm;
static ggml_cl_buffer buf_trans_s;
buf_trans_q.allocate(backend_ctx->context, size_q); buf_trans_q.allocate(backend_ctx->context, size_q);
buf_trans_d.allocate(backend_ctx->context, size_d); buf_trans_d.allocate(backend_ctx->context, size_d);
buf_trans_dm.allocate(backend_ctx->context, size_dm); buf_trans_dm.allocate(backend_ctx->context, size_dm);
buf_trans_s.allocate(backend_ctx->context, size_s);
// Transpose q, d, dm back // Transpose q, d, dm, s back
transpose_2d_as_16b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K/4); transpose_2d_as_16b(backend_ctx, extra->q, buf_trans_q.buffer, size_q, M, K/4);
transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256); transpose_2d_as_16b(backend_ctx, extra->d, buf_trans_d.buffer, size_d, M, K/256);
transpose_2d_as_16b(backend_ctx, extra->dm, buf_trans_dm.buffer, size_dm, M, K/256); transpose_2d_as_16b(backend_ctx, extra->dm, buf_trans_dm.buffer, size_dm, M, K/256);
transpose_2d_as_8b (backend_ctx, extra->s, buf_trans_s.buffer, size_s, M, K/256*12, true, true);
cl_kernel kernel = backend_ctx->kernel_restore_block_q4_K_noshuffle; cl_kernel kernel = backend_ctx->kernel_restore_block_q4_K_noshuffle;
CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &buf_trans_q.buffer)); CL_CHECK(clSetKernelArg(kernel, 0, sizeof(cl_mem), &buf_trans_q.buffer));
CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &extra->s)); CL_CHECK(clSetKernelArg(kernel, 1, sizeof(cl_mem), &buf_trans_s.buffer));
CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &buf_trans_d.buffer)); CL_CHECK(clSetKernelArg(kernel, 2, sizeof(cl_mem), &buf_trans_d.buffer));
CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &buf_trans_dm.buffer)); CL_CHECK(clSetKernelArg(kernel, 3, sizeof(cl_mem), &buf_trans_dm.buffer));
CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &data_device)); CL_CHECK(clSetKernelArg(kernel, 4, sizeof(cl_mem), &data_device));
@@ -8,9 +8,11 @@
#define QK_K 256 #define QK_K 256
#define K_SCALE_SIZE 12 #define K_SCALE_SIZE 12
// scales are transposed: consecutive codes of a row are `stride` apart
inline void get_scale_min_k4( inline void get_scale_min_k4(
int j, int j,
global const uchar * q, global const uchar * q,
int stride,
uchar * d, uchar * d,
uchar * m, uchar * m,
uchar mask_d6, uchar mask_d6,
@@ -18,11 +20,11 @@ inline void get_scale_min_k4(
uchar mask_hi2 uchar mask_hi2
) { ) {
if (j < 4) { if (j < 4) {
*d = q[j] & mask_d6; *d = q[j*stride] & mask_d6;
*m = q[j+4] & mask_d6; *m = q[(j+4)*stride] & mask_d6;
} else { } else {
*d = (q[j+4] & mask_d4) | ((q[j-4] & mask_hi2) >> 2); *d = (q[(j+4)*stride] & mask_d4) | ((q[(j-4)*stride] & mask_hi2) >> 2);
*m = ((q[j+4] >> 4) & mask_d4) | ((q[j] & mask_hi2) >> 2); *m = ((q[(j+4)*stride] >> 4) & mask_d4) | ((q[j*stride] & mask_hi2) >> 2);
} }
} }
@@ -55,7 +57,6 @@ kernel void kernel_gemm_noshuffle_q4_k_f32(
half8 B; half8 B;
half4 dequantized_weights; half4 dequantized_weights;
int num_blocks_K = k / QK_K;
global const ushort * weight_ptr = src0_q + gx_2; global const ushort * weight_ptr = src0_q + gx_2;
global const half * d_ptr = src0_d + gx_2; global const half * d_ptr = src0_d + gx_2;
@@ -68,16 +69,16 @@ kernel void kernel_gemm_noshuffle_q4_k_f32(
half4 d = vload4(0, d_ptr + sb_idx * m); half4 d = vload4(0, d_ptr + sb_idx * m);
half4 dm = vload4(0, dm_ptr + sb_idx * m); half4 dm = vload4(0, dm_ptr + sb_idx * m);
global const uchar * sc0 = src0_s + (gx_2+0) * num_blocks_K * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE; global const uchar * sc0 = src0_s + sb_idx * K_SCALE_SIZE * m + (gx_2+0);
global const uchar * sc1 = src0_s + (gx_2+1) * num_blocks_K * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE; global const uchar * sc1 = sc0 + 1;
global const uchar * sc2 = src0_s + (gx_2+2) * num_blocks_K * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE; global const uchar * sc2 = sc0 + 2;
global const uchar * sc3 = src0_s + (gx_2+3) * num_blocks_K * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE; global const uchar * sc3 = sc0 + 3;
uchar sv0, mn0, sv1, mn1, sv2, mn2, sv3, mn3; uchar sv0, mn0, sv1, mn1, sv2, mn2, sv3, mn3;
get_scale_min_k4(sub_idx, sc0, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(sub_idx, sc0, m, &sv0, &mn0, mask_d6, mask_d4, mask_hi2);
get_scale_min_k4(sub_idx, sc1, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(sub_idx, sc1, m, &sv1, &mn1, mask_d6, mask_d4, mask_hi2);
get_scale_min_k4(sub_idx, sc2, &sv2, &mn2, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(sub_idx, sc2, m, &sv2, &mn2, mask_d6, mask_d4, mask_hi2);
get_scale_min_k4(sub_idx, sc3, &sv3, &mn3, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(sub_idx, sc3, m, &sv3, &mn3, mask_d6, mask_d4, mask_hi2);
half4 scale = convert_half4(convert_float4(d) * convert_float4((uchar4)(sv0, sv1, sv2, sv3))); half4 scale = convert_half4(convert_float4(d) * convert_float4((uchar4)(sv0, sv1, sv2, sv3)));
half4 mval = convert_half4(convert_float4(dm) * convert_float4((uchar4)(mn0, mn1, mn2, mn3))); half4 mval = convert_half4(convert_float4(dm) * convert_float4((uchar4)(mn0, mn1, mn2, mn3)));
@@ -10,9 +10,11 @@
#define QK_K 256 #define QK_K 256
#define K_SCALE_SIZE 12 #define K_SCALE_SIZE 12
// scales are transposed: consecutive codes of a row are `stride` apart
inline void get_scale_min_k4( inline void get_scale_min_k4(
int j, int j,
global const uchar * q, global const uchar * q,
uint stride,
uchar * d, uchar * d,
uchar * m, uchar * m,
uchar mask_d6, uchar mask_d6,
@@ -20,11 +22,11 @@ inline void get_scale_min_k4(
uchar mask_hi2 uchar mask_hi2
) { ) {
if (j < 4) { if (j < 4) {
*d = q[j] & mask_d6; *d = q[j*stride] & mask_d6;
*m = q[j+4] & mask_d6; *m = q[(j+4)*stride] & mask_d6;
} else { } else {
*d = (q[j+4] & mask_d4) | ((q[j-4] & mask_hi2) >> 2); *d = (q[(j+4)*stride] & mask_d4) | ((q[(j-4)*stride] & mask_hi2) >> 2);
*m = ((q[j+4] >> 4) & mask_d4) | ((q[j] & mask_hi2) >> 2); *m = ((q[(j+4)*stride] >> 4) & mask_d4) | ((q[j*stride] & mask_hi2) >> 2);
} }
} }
@@ -79,7 +81,6 @@ kernel void kernel_gemm_noshuffle_q4_k_q8_1_dp4a(
const bool row_valid = row < (uint)m; const bool row_valid = row < (uint)m;
const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked const uint rrow = row_valid ? row : 0; // clamp OOB rows; their writes are masked
const uint num_superblocks = (uint)k / QK_K;
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K const uint k_b = (uint)k >> 5; // blocks-of-32 along K
@@ -101,9 +102,9 @@ kernel void kernel_gemm_noshuffle_q4_k_q8_1_dp4a(
// weight scale/min for this WI's row, this subblock // weight scale/min for this WI's row, this subblock
const float dd = (float)src0_d [rrow + sb_idx * m]; const float dd = (float)src0_d [rrow + sb_idx * m];
const float dmm = (float)src0_dm[rrow + sb_idx * m]; const float dmm = (float)src0_dm[rrow + sb_idx * m];
global const uchar * sc = src0_s + rrow * num_superblocks * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE; global const uchar * sc = src0_s + sb_idx * K_SCALE_SIZE * (uint)m + rrow;
uchar sv, mn; uchar sv, mn;
get_scale_min_k4(sub_idx, sc, &sv, &mn, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(sub_idx, sc, (uint)m, &sv, &mn, mask_d6, mask_d4, mask_hi2);
const float scale = dd * (float)sv; const float scale = dd * (float)sv;
const float minv = dmm * (float)mn; const float minv = dmm * (float)mn;
@@ -202,7 +203,6 @@ kernel void kernel_gemm_noshuffle_q4_k_q8_1_dp4a_wimg(
const uint k_u = (uint)k >> 2; // K in uint (int8x4) units const uint k_u = (uint)k >> 2; // K in uint (int8x4) units
const uint k_b = (uint)k >> 5; // blocks-of-32 along K const uint k_b = (uint)k >> 5; // blocks-of-32 along K
const uint num_superblocks = (uint)k / QK_K;
__local uint sh_qa[TILESIZE_N][8]; __local uint sh_qa[TILESIZE_N][8];
__local half sh_d[TILESIZE_N]; __local half sh_d[TILESIZE_N];
@@ -220,9 +220,9 @@ kernel void kernel_gemm_noshuffle_q4_k_q8_1_dp4a_wimg(
const float dd = (float)src0_d [rrow + sb_idx * m]; const float dd = (float)src0_d [rrow + sb_idx * m];
const float dmm = (float)src0_dm[rrow + sb_idx * m]; const float dmm = (float)src0_dm[rrow + sb_idx * m];
global const uchar * sc = src0_s + rrow * num_superblocks * K_SCALE_SIZE + sb_idx * K_SCALE_SIZE; global const uchar * sc = src0_s + sb_idx * K_SCALE_SIZE * (uint)m + rrow;
uchar sv, mn; uchar sv, mn;
get_scale_min_k4(sub_idx, sc, &sv, &mn, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(sub_idx, sc, (uint)m, &sv, &mn, mask_d6, mask_d4, mask_hi2);
const float scale = dd * (float)sv; const float scale = dd * (float)sv;
const float minv = dmm * (float)mn; const float minv = dmm * (float)mn;
@@ -11,9 +11,11 @@
#define NSUBGROUPS 4 #define NSUBGROUPS 4
#define SUBGROUP_SIZE 64 #define SUBGROUP_SIZE 64
// scales are transposed: consecutive codes of a row are `stride` apart
inline void get_scale_min_k4( inline void get_scale_min_k4(
int j, int j,
global const uchar * q, global const uchar * q,
uint stride,
uchar * d, uchar * d,
uchar * m, uchar * m,
uchar mask_d6, uchar mask_d6,
@@ -21,11 +23,11 @@ inline void get_scale_min_k4(
uchar mask_hi2 uchar mask_hi2
) { ) {
if (j < 4) { if (j < 4) {
*d = q[j] & mask_d6; *d = q[j*stride] & mask_d6;
*m = q[j+4] & mask_d6; *m = q[(j+4)*stride] & mask_d6;
} else { } else {
*d = (q[j+4] & mask_d4) | ((q[j-4] & mask_hi2) >> 2); *d = (q[(j+4)*stride] & mask_d4) | ((q[(j-4)*stride] & mask_hi2) >> 2);
*m = ((q[j+4] >> 4) & mask_d4) | ((q[j] & mask_hi2) >> 2); *m = ((q[(j+4)*stride] >> 4) & mask_d4) | ((q[j*stride] & mask_hi2) >> 2);
} }
} }
@@ -232,7 +234,6 @@ kernel void kernel_gemv_noshuffle_q4_k_f32(
uint LINE_STRIDE_A = M / 2; uint LINE_STRIDE_A = M / 2;
uint BLOCK_STRIDE_A = NSUBGROUPS * M; uint BLOCK_STRIDE_A = NSUBGROUPS * M;
uint scales_per_row = (K / QK_K) * 12;
private uint4 regA; private uint4 regA;
private half2 regS; private half2 regS;
@@ -248,12 +249,12 @@ kernel void kernel_gemv_noshuffle_q4_k_f32(
half2 d = src0_d[gid + sb * LINE_STRIDE_A]; half2 d = src0_d[gid + sb * LINE_STRIDE_A];
half2 dm = src0_m[gid + sb * LINE_STRIDE_A]; half2 dm = src0_m[gid + sb * LINE_STRIDE_A];
global const uchar * sc0 = src0_s + 2 * gid * scales_per_row + sb * 12; global const uchar * sc0 = src0_s + sb * 12 * M + 2 * gid;
global const uchar * sc1 = src0_s + (2 * gid + 1) * scales_per_row + sb * 12; global const uchar * sc1 = sc0 + 1;
uchar sv0, mn0, sv1, mn1; uchar sv0, mn0, sv1, mn1;
get_scale_min_k4(j, sc0, &sv0, &mn0, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(j, sc0, M, &sv0, &mn0, mask_d6, mask_d4, mask_hi2);
get_scale_min_k4(j, sc1, &sv1, &mn1, mask_d6, mask_d4, mask_hi2); get_scale_min_k4(j, sc1, M, &sv1, &mn1, mask_d6, mask_d4, mask_hi2);
regS = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1))); regS = convert_half2(convert_float2(d) * convert_float2((uchar2)(sv0, sv1)));
regM = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1))); regM = convert_half2(convert_float2(dm) * convert_float2((uchar2)(mn0, mn1)));