mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-09-19 09:15:04 +02:00
skip computation for inactive tiles
This commit is contained in:
@@ -334,6 +334,12 @@ void main() {
|
||||
PREFETCH_BLOCK(start_k)
|
||||
}
|
||||
|
||||
const uint a_row0 = warp_r * WM;
|
||||
const uint b_col0 = warp_c * WN;
|
||||
const bool active_col_tile = ic * BN + b_col0 < p.N;
|
||||
|
||||
barrier();
|
||||
|
||||
for (uint block = start_k; block < end_k; block += BK * BK_STEP) {
|
||||
// Store prefetched data to shmem
|
||||
STORE_BLOCK_TO_LDS(block)
|
||||
@@ -349,6 +355,7 @@ void main() {
|
||||
PREFETCH_BLOCK(next_block)
|
||||
}
|
||||
|
||||
if (active_col_tile) {
|
||||
// Compute from shmem
|
||||
[[unroll]] for (uint ks = 0; ks < BK_STEP; ks++) {
|
||||
const uint K_SUB = BK / TK;
|
||||
@@ -357,12 +364,12 @@ void main() {
|
||||
|
||||
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
|
||||
[[unroll]] for (uint h = 0; h < K_SUB; h++) {
|
||||
coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (warp_r * WM + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
|
||||
coopMatLoad(cache_a[r * K_SUB + h], buf_a_qs, (a_row0 + r * TM) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutRowMajor);
|
||||
}
|
||||
}
|
||||
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
|
||||
[[unroll]] for (uint h = 0; h < K_SUB; h++) {
|
||||
coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (warp_c * WN + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
|
||||
coopMatLoad(cache_b[c * K_SUB + h], buf_b_qs, (b_col0 + c * TN) * QPITCH + ks * (BK / 4) + h * (TK / 4), QPITCH, gl_CooperativeMatrixLayoutColumnMajor);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -371,7 +378,7 @@ void main() {
|
||||
float scale_b[cms_per_col * CM_ELEMS];
|
||||
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
scale_a[r * CM_ELEMS + e] = buf_a_d[ks * BM + warp_r * WM + r * TM + elem_row[e]];
|
||||
scale_a[r * CM_ELEMS + e] = buf_a_d[ks * BM + a_row0 + r * TM + elem_row[e]];
|
||||
if (USE_MAGIC_BIAS) {
|
||||
nbias_a[r * CM_ELEMS + e] = -ACC_BIAS_F * scale_a[r * CM_ELEMS + e];
|
||||
}
|
||||
@@ -379,7 +386,7 @@ void main() {
|
||||
}
|
||||
[[unroll]] for (uint c = 0; c < cms_per_col; c++) {
|
||||
[[unroll]] for (uint e = 0; e < CM_ELEMS; e++) {
|
||||
scale_b[c * CM_ELEMS + e] = buf_b_d[ks * BN + warp_c * WN + c * TN + elem_col[e]];
|
||||
scale_b[c * CM_ELEMS + e] = buf_b_d[ks * BN + b_col0 + c * TN + elem_col[e]];
|
||||
}
|
||||
}
|
||||
|
||||
@@ -417,6 +424,7 @@ void main() {
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
barrier();
|
||||
}
|
||||
@@ -426,8 +434,8 @@ void main() {
|
||||
#undef B_IB_CALC
|
||||
#undef PREFETCH_BLOCK
|
||||
|
||||
const uint dr = ir * BM + warp_r * WM;
|
||||
const uint dc = ic * BN + warp_c * WN;
|
||||
const uint dr = ir * BM + a_row0;
|
||||
const uint dc = ic * BN + b_col0;
|
||||
|
||||
#ifdef MUL_MAT_ID
|
||||
[[unroll]] for (uint r = 0; r < cms_per_row; r++) {
|
||||
|
||||
Reference in New Issue
Block a user