preload scales

This commit is contained in:
Ruben Ortlam
2026-08-24 10:37:12 +02:00
parent 21fb8eeef2
commit ab4e443f0f
@@ -360,14 +360,27 @@ void main() {
}
}
// Apply scales directly from coopmat elements
// Pre-load scales into registers
ACC_TYPE scale_a[cms_per_row * CM_ELEMS];
ACC_TYPE 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] = ACC_TYPE(buf_a_d[ks * BM + warp_r * WM + r * TM + elem_row[e]]);
}
}
[[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] = ACC_TYPE(buf_b_d[ks * BN + warp_c * WN + c * TN + elem_col[e]]);
}
}
// Apply scales from registers
[[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;
sums[tile_idx * CM_ELEMS + e] += ACC_TYPE(cm_result[tile_idx][e])
* scale_a[cm_row * CM_ELEMS + e] * scale_b[cm_col * CM_ELEMS + e];
}
}
}