@@ -16,15 +16,11 @@ layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
1616
1717uint a_offset, b_offset, d_offset, y_offset;
1818
19- vec4 load_b(const uint j, const uint iybs, const uint iqs, const bool lastiter) {
20- // Check if the second of the pair of elements is OOB, and don't fetch B or
21- // accumulate it. We still fetch a pair of elements for A, which is fine for
22- // quantized formats since they'll be within the same block. We should
23- // probably skip fetching the second element for F16/F32, but as of now we
24- // still do.
25- const bool OOB_y = lastiter && (iybs + iqs + y_offset >= p.ncols);
26- const bool OOB_z = lastiter && (iybs + iqs + y_offset*2 >= p.ncols);
27- const bool OOB_w = lastiter && (iybs + iqs + y_offset*3 >= p.ncols);
19+ vec4 load_b(const uint j, const uint iybs, const uint iqs, const bool lastiter, out bool OOB_y, out bool OOB_z, out bool OOB_w) {
20+ // Check if the latter elements are OOB, and don't fetch B or accumulate it.
21+ OOB_y = lastiter && (iybs + iqs + y_offset >= p.ncols);
22+ OOB_z = lastiter && (iybs + iqs + y_offset*2 >= p.ncols);
23+ OOB_w = lastiter && (iybs + iqs + y_offset*3 >= p.ncols);
2824
2925 if (!OOB_w) {
3026 return vec4(FLOAT_TYPE(data_b[j*p.batch_stride_b + b_offset + iybs + iqs]),
@@ -46,14 +42,6 @@ vec4 load_b(const uint j, const uint iybs, const uint iqs, const bool lastiter)
4642 }
4743}
4844
49- vec4 load_b_aligned(const uint j, const uint iybs, const uint iqs, const bool lastiter) {
50- const bool OOB_w = lastiter && (iybs + iqs + y_offset*3 >= p.ncols);
51- if (!OOB_w)
52- return data_b_v4[(j*p.batch_stride_b + b_offset + iybs + iqs) / 4];
53- else
54- return load_b(j, iybs, iqs, lastiter);
55- }
56-
5745void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint tid, const uint i, bool lastiter)
5846{
5947 [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
@@ -63,6 +51,8 @@ void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const
6351
6452#if K_PER_ITER == 8
6553#if QUANT_R == 2
54+ // Note that we end up fetching bogus elements here, but its fine as they'll be
55+ // within an accessible block.
6656 const vec4 bv02 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + iybs + iqs) / 4]);
6757 const vec4 bv13 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + iybs + iqs + y_offset) / 4]);
6858 const vec4 bv0 = vec4(bv02.x, bv13.x, bv02.y, bv13.y);
@@ -72,7 +62,11 @@ void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const
7262 const vec4 bv1 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + iybs + iqs) / 4 + 1]);
7363#endif
7464#else
75- const vec4 b = load_b(j, iybs, iqs, lastiter);
65+ bool OOB_y;
66+ bool OOB_z;
67+ bool OOB_w;
68+
69+ const vec4 b = load_b(j, iybs, iqs, lastiter, OOB_y, OOB_z, OOB_w);
7670#endif
7771 uint ibi = first_row*p.ncols;
7872 [[unroll]] for (uint n = 0; n < num_rows; ++n) {
@@ -98,24 +92,38 @@ void iter(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const
9892
9993 temp[j][n] += rowtmp;
10094#else
101- const vec4 v = dequantize4(ib, iqs, a_offset);
102-
103- // matrix multiplication
104- temp[j][n] += dot(v, b);
95+ if (!OOB_w) {
96+ const vec4 v = dequantize4(ib, iqs, a_offset);
97+ temp[j][n] += dot(v, b);
98+ } else if (!OOB_z) {
99+ const vec2 v0 = dequantize(ib, iqs, a_offset);
100+ const FLOAT_TYPE v1 = dequantize1(ib + 2/QUANT_R, iqs, a_offset);
101+ const vec3 v = vec3(v0.x, v0.y, v1);
102+ const vec3 b0 = vec3(b.x, b.y, b.z);
103+ temp[j][n] += dot(v, b0);
104+ } else if (!OOB_y) {
105+ const vec2 v0 = dequantize(ib, iqs, a_offset);
106+ const vec2 b0 = vec2(b.x, b.y);
107+ temp[j][n] += dot(v0, b0);
108+ } else {
109+ const FLOAT_TYPE v = dequantize1(ib, iqs, a_offset);
110+ temp[j][n] = fma(v, b.x, temp[j][n]);
111+ }
105112#endif
106113 }
107114 }
108115}
109116
110117#if defined(DATA_A_F32) || defined(DATA_A_F16) || defined(DATA_A_BF16)
111- void iter_aligned_nonquant(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint tid, const uint i, bool lastiter )
118+ void iter_aligned_nonquant(inout FLOAT_TYPE temp[NUM_COLS][NUM_ROWS], const uint first_row, const uint num_rows, const uint tid, const uint i)
112119{
113120 [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
114121 const uint col = i*BLOCK_SIZE + K_PER_ITER*tid;
115122 const uint iqs = 0; // quant index
116123 const uint iybs = col; // y block start index
117124
118- const vec4 b = load_b_aligned(j, iybs, iqs, lastiter);
125+ const vec4 b = data_b_v4[(j*p.batch_stride_b + b_offset + iybs + iqs) / 4];
126+
119127 uint ibi = first_row*p.ncols;
120128 [[unroll]] for (uint n = 0; n < num_rows; ++n) {
121129 const uint ib = (ibi + col)/QUANT_K; // block index
@@ -136,7 +144,7 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
136144 get_offsets(a_offset, b_offset, d_offset);
137145 const bool is_aligned_nonquant =
138146 p.batch_stride_b % 4 == 0 && b_offset % 4 == 0 &&
139- p.ncols % 2 == 0 && BLOCK_SIZE % 4 == 0 &&
147+ p.ncols % 4 == 0 && BLOCK_SIZE % 4 == 0 &&
140148 K_PER_ITER == 4;
141149
142150 y_offset = QUANT_R == 1 ? 1 : QUANT_K/2;
@@ -170,7 +178,7 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
170178 while (i < unrolled_iters) {
171179 // Manually partially unroll the loop
172180 [[unroll]] for (uint k = 0; k < unroll_count; ++k) {
173- iter_aligned_nonquant(temp, first_row, num_rows, tid, i*K_PER_ITER, false );
181+ iter_aligned_nonquant(temp, first_row, num_rows, tid, i*K_PER_ITER);
174182 i++;
175183 }
176184 }
@@ -201,7 +209,7 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
201209 while (i < unrolled_iters && is_aligned_nonquant) {
202210 // Manually partially unroll the loop
203211 [[unroll]] for (uint k = 0; k < unroll_count; ++k) {
204- iter_aligned_nonquant(temp, first_row, num_rows, tid, i*K_PER_ITER, false );
212+ iter_aligned_nonquant(temp, first_row, num_rows, tid, i*K_PER_ITER);
205213 i++;
206214 }
207215 }
@@ -221,7 +229,7 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
221229#if K_PER_ITER == 4
222230 if (is_aligned_nonquant) {
223231 while (i < num_iters) {
224- iter_aligned_nonquant(temp, first_row, num_rows, tid, i*K_PER_ITER, true );
232+ iter_aligned_nonquant(temp, first_row, num_rows, tid, i*K_PER_ITER);
225233 i++;
226234 }
227235 } else {
0 commit comments