mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-19 09:15:04 +02:00
add BK_STEP to shader, default to 2
This commit is contained in:
@@ -80,15 +80,16 @@ layout (constant_id = 9) const uint TK = 16;
|
||||
layout (constant_id = 10) const uint WARP = 32;
|
||||
|
||||
#define BK 32
|
||||
#define BK_STEP 2
|
||||
|
||||
const uint shmem_stride = (BK / 4) + 4;
|
||||
const uint QPITCH = BK_STEP * (BK / 4) + 4;
|
||||
|
||||
// Shared memory cache
|
||||
shared uint32_t buf_a_qs[BM * shmem_stride];
|
||||
shared float16_t buf_a_d[BM];
|
||||
shared uint32_t buf_a_qs[BM * QPITCH];
|
||||
shared float16_t buf_a_d[BM * BK_STEP];
|
||||
|
||||
shared uint32_t buf_b_qs[BN * shmem_stride];
|
||||
shared float16_t buf_b_d[BN];
|
||||
shared uint32_t buf_b_qs[BN * QPITCH];
|
||||
shared float16_t buf_b_d[BN * BK_STEP];
|
||||
|
||||
#define LOAD_VEC_A (4 * QUANT_R)
|
||||
#define LOAD_VEC_B 16
|
||||
@@ -228,58 +229,65 @@ void main() {
|
||||
sums[i] = ACC_TYPE(0.0);
|
||||
}
|
||||
|
||||
for (uint block = start_k; block < end_k; block += BK) {
|
||||
[[unroll]] for (uint l = 0; loadc_a + l < BM; l += loadstride_a) {
|
||||
const uint buf_ib = loadc_a + l;
|
||||
const uint ib = pos_a_ib + buf_ib * p.stride_a / BK;
|
||||
const uint iqs = loadr_a;
|
||||
|
||||
block_a_to_shmem(buf_ib, ib, iqs);
|
||||
}
|
||||
[[unroll]] for (uint l = 0; loadc_b + l < BN; l += loadstride_b) {
|
||||
const uint buf_ib = loadc_b + l;
|
||||
|
||||
for (uint block = start_k; block < end_k; block += BK * BK_STEP) {
|
||||
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
|
||||
const bool k_in_bounds = block + ks * BK < end_k;
|
||||
[[unroll]] for (uint l = 0; loadc_a + l < BM; l += loadstride_a) {
|
||||
const uint buf_ib = loadc_a + l;
|
||||
const uint ib = pos_a_ib + buf_ib * p.stride_a / BK + ks;
|
||||
if (k_in_bounds) {
|
||||
block_a_to_shmem(buf_ib, ib, loadr_a, ks);
|
||||
} else if (loadr_a == 0) {
|
||||
buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(0.0);
|
||||
}
|
||||
}
|
||||
[[unroll]] for (uint l = 0; loadc_b + l < BN; l += loadstride_b) {
|
||||
const uint buf_ib = loadc_b + l;
|
||||
#ifdef MUL_MAT_ID
|
||||
const u16vec2 row_idx = row_ids[buf_ib];
|
||||
const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK + (row_idx.x % p.ne11) * p.stride_b / BK;
|
||||
const u16vec2 row_idx = row_ids[buf_ib];
|
||||
const uint ib = pos_b_ib + row_idx.y * p.batch_stride_b / BK + (row_idx.x % p.ne11) * p.stride_b / BK + ks;
|
||||
#else
|
||||
const uint ib = pos_b_ib + buf_ib * p.stride_b / BK;
|
||||
const uint ib = pos_b_ib + buf_ib * p.stride_b / BK + ks;
|
||||
#endif
|
||||
const uint iqs = loadr_b;
|
||||
|
||||
block_b_to_shmem(buf_ib, ib, iqs);
|
||||
}
|
||||
|
||||
barrier();
|
||||
|
||||
pos_a_ib += 1;
|
||||
pos_b_ib += 1;
|
||||
|
||||
[[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) {
|
||||
cm_result[idx] = coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0);
|
||||
}
|
||||
|
||||
// Calculate quants
|
||||
[[unroll]] for (uint i = 0; i < BK; i += TK) {
|
||||
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
|
||||
coopMatLoad(cache_a, buf_a_qs, (warp_r * WM + cm_row * TM) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutRowMajor);
|
||||
|
||||
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
|
||||
coopMatLoad(cache_b, buf_b_qs, (warp_c * WN + cm_col * TN) * shmem_stride + i / 4, shmem_stride, gl_CooperativeMatrixLayoutColumnMajor);
|
||||
|
||||
cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]);
|
||||
if (k_in_bounds) {
|
||||
block_b_to_shmem(buf_ib, ib, loadr_b, ks);
|
||||
} else if (loadr_b == 0) {
|
||||
buf_b_d[ks * BN + buf_ib] = FLOAT_TYPE(0.0);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Apply scales directly from coopmat elements
|
||||
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
|
||||
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
|
||||
const uint tile_idx = cm_col * cms_per_row + cm_row;
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
const ACC_TYPE da = ACC_TYPE(buf_a_d[warp_r * WM + cm_row * TM + elem_row[e]]);
|
||||
const ACC_TYPE db = ACC_TYPE(buf_b_d[warp_c * WN + cm_col * TN + elem_col[e]]);
|
||||
sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) * da * db;
|
||||
barrier();
|
||||
|
||||
pos_a_ib += BK_STEP;
|
||||
pos_b_ib += BK_STEP;
|
||||
|
||||
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
|
||||
[[unroll]] for (uint idx = 0; idx < cms_per_row * cms_per_col; idx++) {
|
||||
cm_result[idx] = coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(0);
|
||||
}
|
||||
|
||||
[[unroll]] for (uint i = 0; i < BK; i += TK) {
|
||||
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
|
||||
coopMatLoad(cache_a, buf_a_qs, (warp_r * WM + cm_row * TM) * QPITCH + ks * (BK / 4) + i / 4, QPITCH, gl_CooperativeMatrixLayoutRowMajor);
|
||||
|
||||
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
|
||||
coopMatLoad(cache_b, buf_b_qs, (warp_c * WN + cm_col * TN) * QPITCH + ks * (BK / 4) + i / 4, QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
|
||||
|
||||
cm_result[cm_col * cms_per_row + cm_row] = coopMatMulAdd(cache_a, cache_b, cm_result[cm_col * cms_per_row + cm_row]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Apply scales directly from coopmat elements
|
||||
[[unroll]] for (uint cm_row = 0; cm_row < cms_per_row; cm_row++) {
|
||||
[[unroll]] for (uint cm_col = 0; cm_col < cms_per_col; cm_col++) {
|
||||
const uint tile_idx = cm_col * cms_per_row + cm_row;
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
const ACC_TYPE da = ACC_TYPE(buf_a_d[ks * BM + warp_r * WM + cm_row * TM + elem_row[e]]);
|
||||
const ACC_TYPE db = ACC_TYPE(buf_b_d[ks * BN + warp_c * WN + cm_col * TN + elem_col[e]]);
|
||||
sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e]) * da * db;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@
|
||||
#if defined(DATA_A_Q4_0) || defined(DATA_A_Q4_1)
|
||||
// 2-byte loads for Q4_0 blocks (18 bytes)
|
||||
// 4-byte loads for Q4_1 blocks (20 bytes)
|
||||
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
|
||||
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) {
|
||||
#ifdef DATA_A_Q4_0
|
||||
const uint32_t vui = pack32(u16vec2(data_a_packed16[ib].qs[iqs * 2],
|
||||
data_a_packed16[ib].qs[iqs * 2 + 1]));
|
||||
@@ -24,12 +24,12 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
|
||||
lo4 = ((lo4 | 0x80808080) - 0x08080808) ^ 0x80808080;
|
||||
hi4 = ((hi4 | 0x80808080) - 0x08080808) ^ 0x80808080;
|
||||
|
||||
buf_a_qs[buf_ib * shmem_stride + iqs ] = lo4;
|
||||
buf_a_qs[buf_ib * shmem_stride + iqs + 4] = hi4;
|
||||
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs ] = lo4;
|
||||
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs + 4] = hi4;
|
||||
|
||||
if (iqs == 0) {
|
||||
#ifdef DATA_A_Q4_0
|
||||
buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
|
||||
buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
|
||||
#else // DATA_A_Q4_1
|
||||
#endif
|
||||
}
|
||||
@@ -44,14 +44,14 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
|
||||
|
||||
#if defined(DATA_A_Q8_0)
|
||||
// 2-byte loads for Q8_0 blocks (34 bytes)
|
||||
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
|
||||
void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) {
|
||||
const uint32_t vui = pack32(u16vec2(data_a_packed16[ib].qs[iqs * 2],
|
||||
data_a_packed16[ib].qs[iqs * 2 + 1]));
|
||||
|
||||
buf_a_qs[buf_ib * shmem_stride + iqs] = vui;
|
||||
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs] = vui;
|
||||
|
||||
if (iqs == 0) {
|
||||
buf_a_d[buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
|
||||
buf_a_d[ks * BM + buf_ib] = FLOAT_TYPE(data_a_packed16[ib].d);
|
||||
}
|
||||
}
|
||||
#endif
|
||||
@@ -78,18 +78,17 @@ void block_a_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
|
||||
// 2-byte loads for Q6_K blocks (210 bytes)
|
||||
#endif
|
||||
|
||||
void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs) {
|
||||
void block_b_to_shmem(const uint buf_ib, const uint ib, const uint iqs, const uint ks) {
|
||||
const uint ib_outer = ib / 4;
|
||||
const uint ib_inner = ib % 4;
|
||||
|
||||
if (iqs == 0) {
|
||||
// Divide by TK for matmul scale application
|
||||
buf_b_d[buf_ib] = data_b[ib_outer].ds[ib_inner].x;
|
||||
buf_b_d[ks * BN + buf_ib] = data_b[ib_outer].ds[ib_inner].x;
|
||||
}
|
||||
|
||||
const ivec4 values = data_b[ib_outer].qs[ib_inner * 2 + iqs];
|
||||
buf_b_qs[buf_ib * shmem_stride + iqs * 4 ] = values.x;
|
||||
buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 1] = values.y;
|
||||
buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 2] = values.z;
|
||||
buf_b_qs[buf_ib * shmem_stride + iqs * 4 + 3] = values.w;
|
||||
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 ] = values.x;
|
||||
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 1] = values.y;
|
||||
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 2] = values.z;
|
||||
buf_b_qs[buf_ib * QPITCH + ks * (BK / 4) + iqs * 4 + 3] = values.w;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user