diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index 998cc693c6..42e407ebfe 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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 }; } else { 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; @@ -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, set_rows, true), 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); 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, residual), 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); return; @@ -10911,6 +10921,8 @@ void ggml_vk_rms_norm(ggml_backend_vk_context * ctx, vk_context& subctx, const s std::array elements; 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); ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp index ee813842c0..314c8e565e 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/rms_norm.comp @@ -53,101 +53,107 @@ shared FLOAT_TYPE sumsh[BLOCK_SIZE]; void rms_norm(uint num_iters) { const uint ncols = p.ne00; 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 channel = gl_WorkGroupID.y; - const uint samp = gl_WorkGroupID.z; const uint tid = gl_LocalInvocationID.x; const uint stride_row = p.nb01; const uint stride_channel = p.nb02; const uint stride_sample = p.nb03; - 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(); + // grid.y/z are clamped to the device workgroup limit, iterate over the excess channels/samples + 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 - // Per-row offset in shared memory - uint32_t d_offset = 0; + // Per-row offset in shared memory + uint32_t d_offset = 0; #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 - 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 - 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) { - FLOAT_TYPE xi = FLOAT_TYPE(0); - if (col < ncols) { - xi = FLOAT_TYPE(data_a[a_offset + col]); - } - sum += xi * xi; - } + [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) { + FLOAT_TYPE xi = FLOAT_TYPE(0); + if (col < ncols) { + xi = FLOAT_TYPE(data_a[a_offset + col]); + } + 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; - } - barrier(); - } - sum = sumsh[0]; + // 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; + } + barrier(); + } + sum = sumsh[0]; - const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols); - const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1)); + const FLOAT_TYPE mean = sum / FLOAT_TYPE(ncols); + const FLOAT_TYPE scale = inversesqrt(mean + FLOAT_TYPE(p.param1)); - if (do_multiply) { - if (ncols > p.ne10) { - [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) { - if (col >= ncols) { - continue; - } - FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + fastmod(col, p.ne10)]); + if (do_multiply) { + if (ncols > p.ne10) { + [[unroll]] for (uint col = tid, idx = 0; idx < num_iters; col += BLOCK_SIZE, ++idx) { + if (col >= ncols) { + continue; + } + 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 - value += FLOAT_TYPE(data_c[d_offset + col]); - if (do_post_multiply) { - value *= FLOAT_TYPE(data_e[0]); - } + value += FLOAT_TYPE(data_c[d_offset + col]); + if (do_post_multiply) { + value *= FLOAT_TYPE(data_e[0]); + } #endif - 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; - } - FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]); + 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; + } + FLOAT_TYPE value = scale * FLOAT_TYPE(data_a[a_offset + col]) * FLOAT_TYPE(data_b[b_offset + col]); #if RMS_NORM_ADD_FUSION - value += FLOAT_TYPE(data_c[d_offset + col]); - if (do_post_multiply) { - value *= FLOAT_TYPE(data_e[0]); - } + value += FLOAT_TYPE(data_c[d_offset + col]); + if (do_post_multiply) { + value *= FLOAT_TYPE(data_e[0]); + } #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 - barrier(); - rope_params rp = p.rope; - for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) { - if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) { - rope_neox(t, row, channel, samp, rp); - } else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) { - rope_norm(t, row, channel, samp, rp); + barrier(); + rope_params rp = p.rope; + for (uint t = 2*tid; t < ncols; t += 2*BLOCK_SIZE) { + if (rp.rope_mode == GGML_ROPE_TYPE_NEOX) { + rope_neox(t, row, channel, samp, rp); + } else if (rp.rope_mode == GGML_ROPE_TYPE_NORMAL) { + rope_norm(t, row, channel, samp, rp); + } + } +#endif } } -#endif } void main() {