|
|
@@ -5,23 +5,24 @@
|
|
|
|
|
|
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
|
|
|
|
|
|
-shared FLOAT_TYPE sccache1[BLOCK_SIZE/16][16];
|
|
|
-shared FLOAT_TYPE sccache2[BLOCK_SIZE/16][16];
|
|
|
+shared FLOAT_TYPE sccache1[2][BLOCK_SIZE/16][16];
|
|
|
+shared FLOAT_TYPE sccache2[2][BLOCK_SIZE/16][16];
|
|
|
|
|
|
FLOAT_TYPE temp[NUM_COLS][NUM_ROWS];
|
|
|
+uint csel = 0;
|
|
|
|
|
|
void calc_superblock(const uint a_offset, const uint b_offset, const uint itid, const uint v_im, const uint ix, const uint q_offset, const uint y_offset, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows, const bool all_threads) {
|
|
|
const uint y_idx = i * QUANT_K + y_offset;
|
|
|
|
|
|
[[unroll]] for (uint n = 0; n < num_rows; ++n) {
|
|
|
const uint ib0 = a_offset / QUANT_K + (first_row+n)*num_blocks_per_row;
|
|
|
+ csel ^= 1;
|
|
|
|
|
|
- barrier();
|
|
|
if (!all_threads) { // when we don't have enough blocks to use all threads
|
|
|
if (i < num_blocks_per_row) {
|
|
|
const uint32_t scale = uint32_t(data_a[ib0 + i].scales[itid]);
|
|
|
- sccache1[ix][itid] = FLOAT_TYPE(scale & 0xF);
|
|
|
- sccache2[ix][itid] = FLOAT_TYPE((scale >> 4) & 0xF);
|
|
|
+ sccache1[csel][ix][itid] = FLOAT_TYPE(scale & 0xF);
|
|
|
+ sccache2[csel][ix][itid] = FLOAT_TYPE((scale >> 4) & 0xF);
|
|
|
}
|
|
|
barrier();
|
|
|
|
|
|
@@ -29,8 +30,8 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint itid,
|
|
|
continue;
|
|
|
} else {
|
|
|
const uint32_t scale = uint32_t(data_a[ib0 + i].scales[itid]);
|
|
|
- sccache1[ix][itid] = FLOAT_TYPE(scale & 0xF);
|
|
|
- sccache2[ix][itid] = FLOAT_TYPE((scale >> 4) & 0xF);
|
|
|
+ sccache1[csel][ix][itid] = FLOAT_TYPE(scale & 0xF);
|
|
|
+ sccache2[csel][ix][itid] = FLOAT_TYPE((scale >> 4) & 0xF);
|
|
|
barrier();
|
|
|
}
|
|
|
|
|
|
@@ -57,22 +58,22 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint itid,
|
|
|
FLOAT_TYPE sum1 = FLOAT_TYPE(0.0);
|
|
|
FLOAT_TYPE sum2 = FLOAT_TYPE(0.0);
|
|
|
[[unroll]] for (int l = 0; l < 2; ++l) {
|
|
|
- sum1 = fma(FLOAT_TYPE(b0[l]), sccache1[ix][ 8*v_im] * qs_u32_0[l ],
|
|
|
- fma(FLOAT_TYPE(b16[l]), sccache1[ix][1 + 8*v_im] * qs_u32_0[l+2],
|
|
|
- fma(FLOAT_TYPE(b32[l]), sccache1[ix][2 + 8*v_im] * qs_u32_2[l ],
|
|
|
- fma(FLOAT_TYPE(b48[l]), sccache1[ix][3 + 8*v_im] * qs_u32_2[l+2],
|
|
|
- fma(FLOAT_TYPE(b64[l]), sccache1[ix][4 + 8*v_im] * qs_u32_4[l ],
|
|
|
- fma(FLOAT_TYPE(b80[l]), sccache1[ix][5 + 8*v_im] * qs_u32_4[l+2],
|
|
|
- fma(FLOAT_TYPE(b96[l]), sccache1[ix][6 + 8*v_im] * qs_u32_6[l ],
|
|
|
- fma(FLOAT_TYPE(b112[l]), sccache1[ix][7 + 8*v_im] * qs_u32_6[l+2], sum1))))))));
|
|
|
- sum2 = fma(FLOAT_TYPE(b0[l]), sccache2[ix][ 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b16[l]), sccache2[ix][1 + 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b32[l]), sccache2[ix][2 + 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b48[l]), sccache2[ix][3 + 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b64[l]), sccache2[ix][4 + 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b80[l]), sccache2[ix][5 + 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b96[l]), sccache2[ix][6 + 8*v_im],
|
|
|
- fma(FLOAT_TYPE(b112[l]), sccache2[ix][7 + 8*v_im], sum2))))))));
|
|
|
+ sum1 = fma(FLOAT_TYPE(b0[l]), sccache1[csel][ix][ 8*v_im] * qs_u32_0[l ],
|
|
|
+ fma(FLOAT_TYPE(b16[l]), sccache1[csel][ix][1 + 8*v_im] * qs_u32_0[l+2],
|
|
|
+ fma(FLOAT_TYPE(b32[l]), sccache1[csel][ix][2 + 8*v_im] * qs_u32_2[l ],
|
|
|
+ fma(FLOAT_TYPE(b48[l]), sccache1[csel][ix][3 + 8*v_im] * qs_u32_2[l+2],
|
|
|
+ fma(FLOAT_TYPE(b64[l]), sccache1[csel][ix][4 + 8*v_im] * qs_u32_4[l ],
|
|
|
+ fma(FLOAT_TYPE(b80[l]), sccache1[csel][ix][5 + 8*v_im] * qs_u32_4[l+2],
|
|
|
+ fma(FLOAT_TYPE(b96[l]), sccache1[csel][ix][6 + 8*v_im] * qs_u32_6[l ],
|
|
|
+ fma(FLOAT_TYPE(b112[l]), sccache1[csel][ix][7 + 8*v_im] * qs_u32_6[l+2], sum1))))))));
|
|
|
+ sum2 = fma(FLOAT_TYPE(b0[l]), sccache2[csel][ix][ 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b16[l]), sccache2[csel][ix][1 + 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b32[l]), sccache2[csel][ix][2 + 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b48[l]), sccache2[csel][ix][3 + 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b64[l]), sccache2[csel][ix][4 + 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b80[l]), sccache2[csel][ix][5 + 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b96[l]), sccache2[csel][ix][6 + 8*v_im],
|
|
|
+ fma(FLOAT_TYPE(b112[l]), sccache2[csel][ix][7 + 8*v_im], sum2))))))));
|
|
|
}
|
|
|
temp[j][n] = fma(dall, sum1, fma(-dmin, sum2, temp[j][n]));
|
|
|
}
|