Skip to content

Commit d853bfe

Browse files
TheBlueMattYour Name
authored andcommitted
vulkan: Fix OOB A reads in MUL_MAT_VEC for odd sizes
There was a TODO to fix the OOB reads from the A matrix which we do here. It is within performance noise (+<0.1%) in tg128 for Qwen3.5-9B:BF16 on Intel BMG.
1 parent f3f4c56 commit d853bfe

2 files changed

Lines changed: 45 additions & 28 deletions

File tree

‎ggml/src/ggml-vulkan/vulkan-shaders/dequant_funcs.glsl‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,9 @@
55
#include "types.glsl"
66

77
#if defined(DATA_A_F32)
8+
FLOAT_TYPE dequantize1(uint ib, uint iqs, uint a_offset) {
9+
return data_a[a_offset + ib];
10+
}
811
vec2 dequantize(uint ib, uint iqs, uint a_offset) {
912
return vec2(data_a[a_offset + ib], data_a[a_offset + ib + 1]);
1013
}
@@ -20,6 +23,9 @@ vec4 dequantize4_2aligned(uint ib, uint iqs, uint a_offset) {
2023
#endif
2124

2225
#if defined(DATA_A_F16)
26+
FLOAT_TYPE dequantize1(uint ib, uint iqs, uint a_offset) {
27+
return data_a[a_offset + ib];
28+
}
2329
vec2 dequantize(uint ib, uint iqs, uint a_offset) {
2430
return vec2(data_a[a_offset + ib], data_a[a_offset + ib + 1]);
2531
}
@@ -35,6 +41,9 @@ vec4 dequantize4_2aligned(uint ib, uint iqs, uint a_offset) {
3541
#endif
3642

3743
#if defined(DATA_A_BF16)
44+
FLOAT_TYPE dequantize1(uint ib, uint iqs, uint a_offset) {
45+
return bf16_to_fp32(data_a[a_offset + ib]);
46+
}
3847
vec2 dequantize(uint ib, uint iqs, uint a_offset) {
3948
return vec2(bf16_to_fp32(data_a[a_offset + ib]), bf16_to_fp32(data_a[a_offset + ib + 1]));
4049
}

‎ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec.comp‎

Lines changed: 36 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -16,15 +16,11 @@ layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
1616

1717
uint 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-
5745
void 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

Comments
 (0)