mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-11 15:30:39 +02:00
vulkan: fix rms_norm workgroup count overflow (#30145)
This commit is contained in:
@@ -9455,6 +9455,8 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co
|
|||||||
elements = { (uint32_t)CEIL_DIV(ne00, 128), 1, 1 };
|
elements = { (uint32_t)CEIL_DIV(ne00, 128), 1, 1 };
|
||||||
} else {
|
} else {
|
||||||
elements = { (uint32_t)ne01, (uint32_t)ne02, (uint32_t)ne03 };
|
elements = { (uint32_t)ne01, (uint32_t)ne02, (uint32_t)ne03 };
|
||||||
|
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
|
||||||
|
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||||
}
|
}
|
||||||
break;
|
break;
|
||||||
|
|
||||||
@@ -10777,7 +10779,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
|||||||
ggml_vk_tensor_subbuffer(ctx, src0, true),
|
ggml_vk_tensor_subbuffer(ctx, src0, true),
|
||||||
ggml_vk_tensor_subbuffer(ctx, set_rows, true),
|
ggml_vk_tensor_subbuffer(ctx, set_rows, true),
|
||||||
ggml_vk_tensor_subbuffer(ctx, indices),
|
ggml_vk_tensor_subbuffer(ctx, indices),
|
||||||
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] });
|
}, pc, {
|
||||||
|
(uint32_t)src0->ne[1],
|
||||||
|
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
|
||||||
|
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
|
||||||
|
});
|
||||||
ggml_vk_rms_norm_finish(ctx, src0);
|
ggml_vk_rms_norm_finish(ctx, src0);
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
@@ -10824,7 +10830,11 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
|||||||
ggml_vk_tensor_subbuffer(ctx, dst, true),
|
ggml_vk_tensor_subbuffer(ctx, dst, true),
|
||||||
ggml_vk_tensor_subbuffer(ctx, residual),
|
ggml_vk_tensor_subbuffer(ctx, residual),
|
||||||
ggml_vk_tensor_subbuffer(ctx, post_scale),
|
ggml_vk_tensor_subbuffer(ctx, post_scale),
|
||||||
}, pc, { (uint32_t)src0->ne[1], (uint32_t)src0->ne[2], (uint32_t)src0->ne[3] });
|
}, pc, {
|
||||||
|
(uint32_t)src0->ne[1],
|
||||||
|
std::min((uint32_t)src0->ne[2], ctx->device->properties.limits.maxComputeWorkGroupCount[1]),
|
||||||
|
std::min((uint32_t)src0->ne[3], ctx->device->properties.limits.maxComputeWorkGroupCount[2]),
|
||||||
|
});
|
||||||
}
|
}
|
||||||
ggml_vk_rms_norm_finish(ctx, src0);
|
ggml_vk_rms_norm_finish(ctx, src0);
|
||||||
return;
|
return;
|
||||||
@@ -10911,6 +10921,8 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s
|
|||||||
|
|
||||||
std::array<uint32_t, 3> elements;
|
std::array<uint32_t, 3> elements;
|
||||||
elements = { (uint32_t)rms->src[0]->ne[1], (uint32_t)rms->src[0]->ne[2], (uint32_t)rms->src[0]->ne[3] };
|
elements = { (uint32_t)rms->src[0]->ne[1], (uint32_t)rms->src[0]->ne[2], (uint32_t)rms->src[0]->ne[3] };
|
||||||
|
elements[1] = std::min(elements[1], ctx->device->properties.limits.maxComputeWorkGroupCount[1]);
|
||||||
|
elements[2] = std::min(elements[2], ctx->device->properties.limits.maxComputeWorkGroupCount[2]);
|
||||||
|
|
||||||
static_assert(max_tensors == 7);
|
static_assert(max_tensors == 7);
|
||||||
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
|
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
|
||||||
|
|||||||
@@ -53,101 +53,107 @@ shared FLOAT_TYPE sumsh[BLOCK_SIZE];
|
|||||||
void rms_norm(uint num_iters) {
|
void rms_norm(uint num_iters) {
|
||||||
const uint ncols = p.ne00;
|
const uint ncols = p.ne00;
|
||||||
const uint nrows = gl_NumWorkGroups.x;
|
const uint nrows = gl_NumWorkGroups.x;
|
||||||
const uint nchannels = gl_NumWorkGroups.y;
|
const uint nchannels = p.ne02;
|
||||||
|
const uint nsamples = p.ne03;
|
||||||
|
|
||||||
const uint row = gl_WorkGroupID.x;
|
const uint row = gl_WorkGroupID.x;
|
||||||
const uint channel = gl_WorkGroupID.y;
|
|
||||||
const uint samp = gl_WorkGroupID.z;
|
|
||||||
const uint tid = gl_LocalInvocationID.x;
|
const uint tid = gl_LocalInvocationID.x;
|
||||||
|
|
||||||
const uint stride_row = p.nb01;
|
const uint stride_row = p.nb01;
|
||||||
const uint stride_channel = p.nb02;
|
const uint stride_channel = p.nb02;
|
||||||
const uint stride_sample = p.nb03;
|
const uint stride_sample = p.nb03;
|
||||||
|
|
||||||
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
|
// grid.y/z are clamped to the device workgroup limit, iterate over the excess channels/samples
|
||||||
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
|
for (uint samp = gl_WorkGroupID.z; samp < nsamples; samp += gl_NumWorkGroups.z) {
|
||||||
|
for (uint channel = gl_WorkGroupID.y; channel < nchannels; channel += gl_NumWorkGroups.y) {
|
||||||
|
barrier();
|
||||||
|
|
||||||
|
uint32_t a_offset = samp*stride_sample + channel*stride_channel + row*stride_row + get_aoffset();
|
||||||
|
uint32_t b_offset = src1_idx(0, row, channel, samp) + get_boffset();
|
||||||
#if RMS_NORM_ROPE_FUSION
|
#if RMS_NORM_ROPE_FUSION
|
||||||
// Per-row offset in shared memory
|
// Per-row offset in shared memory
|
||||||
uint32_t d_offset = 0;
|
uint32_t d_offset = 0;
|
||||||
#elif RMS_NORM_SET_ROWS_FUSION
|
#elif RMS_NORM_SET_ROWS_FUSION
|
||||||
uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
|
uint32_t d_offset = data_i[channel].x*p.nb21 + row*ncols + get_doffset();
|
||||||
#else
|
#else
|
||||||
uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
|
uint32_t d_offset = ((samp*nchannels + channel)*nrows + row)*ncols + get_doffset();
|
||||||
#endif
|
#endif
|
||||||
FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
FLOAT_TYPE sum = FLOAT_TYPE(0.0f); // partial sum for thread in warp
|
||||||
|
|
||||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||||
FLOAT_TYPE xi = FLOAT_TYPE(0);
|
FLOAT_TYPE xi = FLOAT_TYPE(0);
|
||||||
if (col < ncols) {
|
if (col < ncols) {
|
||||||
xi = FLOAT_TYPE(data_a[a_offset + col]);
|
xi = FLOAT_TYPE(data_a[a_offset + col]);
|
||||||
}
|
}
|
||||||
sum += xi * xi;
|
sum += xi * xi;
|
||||||
}
|
}
|
||||||
|
|
||||||
sumsh[tid] = sum;
|
|
||||||
// sum up partial sums and write back result
|
|
||||||
barrier();
|
|
||||||
[[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
|
|
||||||
if (tid < s) {
|
|
||||||
sum += sumsh[tid + s];
|
|
||||||
sumsh[tid] = sum;
|
sumsh[tid] = sum;
|
||||||
}
|
// sum up partial sums and write back result
|
||||||
barrier();
|
barrier();
|
||||||
}
|
[[unroll]] for (int s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
|
||||||
sum = sumsh[0];
|
if (tid < s) {
|
||||||
|
sum += sumsh[tid + s];
|
||||||
|
sumsh[tid] = sum;
|
||||||
|
}
|
||||||
|
barrier();
|
||||||
|
}
|
||||||
|
sum = sumsh[0];
|
||||||
|
|
||||||
const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
|
const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols);
|
||||||
const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
|
const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1));
|
||||||
|
|
||||||
if (do_multiply) {
|
if (do_multiply) {
|
||||||
if (ncols > p.ne10) {
|
if (ncols > p.ne10) {
|
||||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||||
if (col >= ncols) {
|
if (col >= ncols) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
|
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]);
|
||||||
#if RMS_NORM_ADD_FUSION
|
#if RMS_NORM_ADD_FUSION
|
||||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||||
if (do_post_multiply) {
|
if (do_post_multiply) {
|
||||||
value *= FLOAT_TYPE(data_e[0]);
|
value *= FLOAT_TYPE(data_e[0]);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
data_d[d_offset + col] = D_TYPE(value);
|
data_d[d_offset + col] = D_TYPE(value);
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||||
if (col >= ncols) {
|
if (col >= ncols) {
|
||||||
continue;
|
continue;
|
||||||
}
|
}
|
||||||
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
|
FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]);
|
||||||
#if RMS_NORM_ADD_FUSION
|
#if RMS_NORM_ADD_FUSION
|
||||||
value += FLOAT_TYPE(data_c[d_offset + col]);
|
value += FLOAT_TYPE(data_c[d_offset + col]);
|
||||||
if (do_post_multiply) {
|
if (do_post_multiply) {
|
||||||
value *= FLOAT_TYPE(data_e[0]);
|
value *= FLOAT_TYPE(data_e[0]);
|
||||||
}
|
}
|
||||||
#endif
|
#endif
|
||||||
data_d[d_offset + col] = D_TYPE(value);
|
data_d[d_offset + col] = D_TYPE(value);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
} else {
|
||||||
|
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
||||||
|
if (col >= ncols) {
|
||||||
|
continue;
|
||||||
|
}
|
||||||
|
data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
}
|
|
||||||
} else {
|
|
||||||
[[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) {
|
|
||||||
if (col >= ncols) {
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
data_d[d_offset + col] = D_TYPE(scale * FLOAT_TYPE(data_a[a_offset + col]));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
#if RMS_NORM_ROPE_FUSION
|
#if RMS_NORM_ROPE_FUSION
|
||||||
barrier();
|
barrier();
|
||||||
rope_params rp = p.rope;
|
rope_params rp = p.rope;
|
||||||
for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
|
for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) {
|
||||||
if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
|
if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) {
|
||||||
rope_neox(t, row, channel, samp, rp);
|
rope_neox(t, row, channel, samp, rp);
|
||||||
} else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
|
} else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) {
|
||||||
rope_norm(t, row, channel, samp, rp);
|
rope_norm(t, row, channel, samp, rp);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
#endif
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
#endif
|
|
||||||
}
|
}
|
||||||
|
|
||||||
void main() {
|
void main() {
|
||||||
|
|||||||
Reference in New Issue
Block a user