mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-24 13:37:01 +02:00
merge shmem arrays
This commit is contained in:
@@ -95,13 +95,16 @@ const uint QPITCH = BK_STEP * (BK / 4) + 4;
|
||||
|
||||
// Shared memory cache
|
||||
shared uint32_t buf_a_qs[BM * QPITCH];
|
||||
#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
|
||||
shared vec2 buf_a_dm[BM * BK_STEP]; // .x = d (== old buf_a_d), .y = m (== old buf_a_m)
|
||||
#else
|
||||
shared float buf_a_d[BM * BK_STEP * KSCALES];
|
||||
#endif
|
||||
|
||||
shared uint32_t buf_b_qs[BN * QPITCH];
|
||||
shared float buf_b_d[BN * BK_STEP];
|
||||
|
||||
#if defined(DATA_A_Q4_1) || defined(DATA_A_Q5_1) || defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
|
||||
shared float buf_a_m[BM * BK_STEP];
|
||||
shared float buf_b_s[BN * BK_STEP];
|
||||
#endif
|
||||
|
||||
@@ -205,8 +208,8 @@ void main() {
|
||||
const uint warp_c = warp_i / (BM / WM);
|
||||
|
||||
// Probe coopmat element layout to discover row/col mapping
|
||||
uint elem_row[CM_ELEMS];
|
||||
uint elem_col[CM_ELEMS];
|
||||
uint8_t elem_row[CM_ELEMS];
|
||||
uint8_t elem_col[CM_ELEMS];
|
||||
|
||||
for (uint i = gl_LocalInvocationID.x; i < TM * TN; i += BLOCK_SIZE) {
|
||||
cm_layout_probe[i] = int32_t(i);
|
||||
@@ -217,8 +220,8 @@ void main() {
|
||||
coopMatLoad(probe, cm_layout_probe, 0, TN, gl_CooperativeMatrixLayoutRowMajor);
|
||||
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; ++e) {
|
||||
elem_row[e] = uint(probe[e]) / TN;
|
||||
elem_col[e] = uint(probe[e]) % TN;
|
||||
elem_row[e] = uint8_t(uint(probe[e]) / TN);
|
||||
elem_col[e] = uint8_t(uint(probe[e]) % TN);
|
||||
}
|
||||
|
||||
const uint loadr_a = gl_LocalInvocationID.x % (BK / LOAD_VEC_A);
|
||||
@@ -333,7 +336,7 @@ void main() {
|
||||
float scale_a[CM_ELEMS];
|
||||
float nbias_a[CM_ELEMS];
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + elem_row[e]];
|
||||
scale_a[e] = buf_a_d[(ks * KSCALES + h) * BM + a_row0 + r * TM + uint(elem_row[e])];
|
||||
if (USE_MAGIC_BIAS) {
|
||||
nbias_a[e] = -ACC_BIAS_F * scale_a[e];
|
||||
}
|
||||
@@ -343,7 +346,7 @@ void main() {
|
||||
coopmat<int8_t, gl_ScopeSubgroup, TK, TN, gl_MatrixUseB> cache_b;
|
||||
coopMatLoad(cache_b, buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
|
||||
|
||||
const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]];
|
||||
const float scale_b_v = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])];
|
||||
|
||||
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator> acc =
|
||||
coopmat<int32_t, gl_ScopeSubgroup, TM, TN, gl_MatrixUseAccumulator>(
|
||||
@@ -385,15 +388,15 @@ void main() {
|
||||
float nbias_a[CM_ELEMS];
|
||||
float ma[CM_ELEMS];
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]];
|
||||
vec2 dm = buf_a_dm[ks * BM + a_row0 + r * TM + uint(elem_row[e])];
|
||||
scale_a[e] = dm.x;
|
||||
if (USE_MAGIC_BIAS) {
|
||||
nbias_a[e] = -ACC_BIAS_F * scale_a[e];
|
||||
}
|
||||
#if defined(DATA_A_Q4_K) || defined(DATA_A_Q5_K)
|
||||
ma[e] = float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]);
|
||||
ma[e] = dm.y;
|
||||
#else
|
||||
ma[e] = float(QUANT_OFFSET) * scale_a[e]
|
||||
+ float(buf_a_m[ks * BM + a_row0 + r * TM + elem_row[e]]);
|
||||
ma[e] = float(QUANT_OFFSET) * scale_a[e] + dm.y;
|
||||
#endif
|
||||
}
|
||||
|
||||
@@ -407,8 +410,8 @@ void main() {
|
||||
}
|
||||
|
||||
const uint tile_idx = r * cms_per_col + c;
|
||||
const float sb = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]];
|
||||
const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + elem_col[0]]);
|
||||
const float sb = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])];
|
||||
const float bs = float(buf_b_s[ks * BN + b_col0 + c * TN + uint(elem_col[0])]);
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
if (USE_MAGIC_BIAS) {
|
||||
const float t = fma(intBitsToFloat(int(acc[e])),
|
||||
@@ -441,14 +444,14 @@ void main() {
|
||||
|
||||
float scale_b[cms_per_col];
|
||||
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
|
||||
scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[0]];
|
||||
scale_b[c] = buf_b_d[ks * BN + b_col0 + c * TN + uint(elem_col[0])];
|
||||
}
|
||||
|
||||
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
|
||||
float scale_a[CM_ELEMS];
|
||||
float nbias_a[CM_ELEMS];
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]];
|
||||
scale_a[e] = buf_a_d[ks * BM + a_row0 + r * TM + uint(elem_row[e])];
|
||||
if (USE_MAGIC_BIAS) {
|
||||
nbias_a[e] = -ACC_BIAS_F * scale_a[e];
|
||||
}
|
||||
@@ -497,10 +500,10 @@ void main() {
|
||||
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
|
||||
const uint tile_idx = r * cms_per_col + c;
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
const uint col_i = dc + c * TN + elem_col[e];
|
||||
const uint col_i = dc + c * TN + uint(elem_col[e]);
|
||||
if (col_i >= _ne1) continue;
|
||||
|
||||
const uint row_g = dr + r * TM + elem_row[e];
|
||||
const uint row_g = dr + r * TM + uint(elem_row[e]);
|
||||
if (row_g >= p.M) continue;
|
||||
|
||||
const u16vec2 row_idx = row_ids[col_i - ic * BN];
|
||||
@@ -516,8 +519,8 @@ void main() {
|
||||
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
|
||||
const uint tile_idx = r * cms_per_col + c;
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
const uint row_g = dr + r * TM + elem_row[e];
|
||||
const uint col_g = dc + c * TN + elem_col[e];
|
||||
const uint row_g = dr + r * TM + uint(elem_row[e]);
|
||||
const uint col_g = dc + c * TN + uint(elem_col[e]);
|
||||
if (row_g < p.M && col_g < p.N) {
|
||||
data_d[offsets + col_g * p.stride_d + row_g] = D_TYPE(sums[tile_idx * CM_ELEMS + e]);
|
||||
}
|
||||
|
||||
@@ -55,8 +55,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
|
||||
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
|
||||
|
||||
if (loadr == 0) {
|
||||
buf_a_d[ks * BM + buf_ib] = float(blk.dm.x);
|
||||
buf_a_m[ks * BM + buf_ib] = float(blk.dm.y);
|
||||
buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -119,8 +118,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
|
||||
buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] = hi4;
|
||||
|
||||
if (loadr == 0) {
|
||||
buf_a_d[ks * BM + buf_ib] = float(blk.dm.x);
|
||||
buf_a_m[ks * BM + buf_ib] = float(blk.dm.y);
|
||||
buf_a_dm[ks * BM + buf_ib] = vec2(float(blk.dm.x), float(blk.dm.y));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -258,8 +256,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
|
||||
}
|
||||
vec2 dm = vec2(data_a_packed32[ib_k].dm);
|
||||
float d_scaled = dm.x * float(sc_val);
|
||||
buf_a_d[ks * BM + buf_ib] = d_scaled;
|
||||
buf_a_m[ks * BM + buf_ib] = 8.0 * d_scaled - (dm.y * float(mn_val));
|
||||
buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 8.0 * d_scaled - (dm.y * float(mn_val)));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -316,8 +313,7 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
|
||||
}
|
||||
vec2 dm = vec2(data_a_packed32[ib_k].dm);
|
||||
float d_scaled = dm.x * float(sc_val);
|
||||
buf_a_d[ks * BM + buf_ib] = d_scaled;
|
||||
buf_a_m[ks * BM + buf_ib] = 16.0 * d_scaled - (dm.y * float(mn_val));
|
||||
buf_a_dm[ks * BM + buf_ib] = vec2(d_scaled, 16.0 * d_scaled - (dm.y * float(mn_val)));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user