vulkan: support qwen4exp hc ops (#28988)

* vulkan: support qwen4exp hc ops

* fix stale comment [no-ci]
This commit is contained in:
Ruben Ortlam
2026-09-17 06:34:23 +02:00
committed by GitHub
parent aa39d7a3e1
commit 35822afe58
3 changed files with 62 additions and 23 deletions
+25 -11
View File
@@ -1130,7 +1130,9 @@ struct vk_device_struct {
vk_pipeline pipeline_count_equal_i32;
vk_pipeline pipeline_dsv4_hc_comb_f32;
vk_pipeline pipeline_dsv4_hc_pre_f32;
vk_pipeline pipeline_dsv4_hc_pre_gated_f32;
vk_pipeline pipeline_dsv4_hc_post_f32;
vk_pipeline pipeline_dsv4_hc_post_nocomb_f32;
std::map<vk_solve_tri_pipeline_state, vk_pipeline> pipeline_solve_tri_f32;
vk_pipeline pipeline_im2col_f32, pipeline_im2col_f32_f16;
vk_pipeline pipeline_im2col_3d_f32, pipeline_im2col_3d_f32_f16;
@@ -1514,12 +1516,14 @@ struct vk_op_dsv4_hc_pre_push_constants {
uint32_t n_tokens;
uint32_t nbx0; uint32_t nbx1; uint32_t nbx2;
uint32_t nbw0; uint32_t nbw1;
uint32_t nbw0; uint32_t nbw1; uint32_t nbw2;
uint32_t nbd0; uint32_t nbd1;
uint32_t x_offset;
uint32_t w_offset;
uint32_t d_offset;
float scale;
};
struct vk_op_dsv4_hc_post_push_constants {
@@ -2737,7 +2741,7 @@ template <> void init_pushconst_tensor_offsets(ggml_backend_vk_context * ctx, vk
p.x_offset = get_misalign_bytes(ctx, src0) / ggml_type_size(src0->type);
p.r_offset = get_misalign_bytes(ctx, src1) / ggml_type_size(src1->type);
p.p_offset = get_misalign_bytes(ctx, src2) / ggml_type_size(src2->type);
p.c_offset = get_misalign_bytes(ctx, src3) / ggml_type_size(src3->type);
p.c_offset = src3 ? get_misalign_bytes(ctx, src3) / ggml_type_size(src3->type) : 0;
p.d_offset = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type);
}
@@ -6173,8 +6177,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_comb_f32, "dsv4_hc_comb_f32", dsv4_hc_comb_f32_len, dsv4_hc_comb_f32_data, "main", 4, sizeof(vk_op_dsv4_hc_comb_push_constants), {tokens_per_workgroup, 1, 1}, { device->subgroup_size }, 1, true, true, device->subgroup_size);
}
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_pre_f32, "dsv4_hc_pre_f32", dsv4_hc_pre_f32_len, dsv4_hc_pre_f32_data, "main", 3, sizeof(vk_op_dsv4_hc_pre_push_constants), {256, 1, 1}, { 256 }, 1);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_post_f32, "dsv4_hc_post_f32", dsv4_hc_post_f32_len, dsv4_hc_post_f32_data, "main", 5, sizeof(vk_op_dsv4_hc_post_push_constants), {256, 1, 1}, { 256 }, 1);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_pre_f32, "dsv4_hc_pre_f32", dsv4_hc_pre_f32_len, dsv4_hc_pre_f32_data, "main", 3, sizeof(vk_op_dsv4_hc_pre_push_constants), {256, 1, 1}, { 256, 0 }, 1);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_pre_gated_f32, "dsv4_hc_pre_gated_f32", dsv4_hc_pre_f32_len, dsv4_hc_pre_f32_data, "main", 3, sizeof(vk_op_dsv4_hc_pre_push_constants), {256, 1, 1}, { 256, 1 }, 1);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_post_f32, "dsv4_hc_post_f32", dsv4_hc_post_f32_len, dsv4_hc_post_f32_data, "main", 5, sizeof(vk_op_dsv4_hc_post_push_constants), {256, 1, 1}, { 256, 1 }, 1);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_post_nocomb_f32,"dsv4_hc_post_nocomb_f32",dsv4_hc_post_f32_len, dsv4_hc_post_f32_data, "main", 5, sizeof(vk_op_dsv4_hc_post_push_constants), {256, 1, 1}, { 256, 0 }, 1);
for (auto &s : device->pipeline_solve_tri_f32) {
const vk_solve_tri_pipeline_state &state = s.first;
@@ -10356,7 +10362,10 @@ static void ggml_vk_dsv4_hc_comb(ggml_backend_vk_context * ctx, vk_context& subc
static void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * weights, ggml_tensor * dst) {
VK_LOG_DEBUG("ggml_vk_dsv4_hc_pre(" << x << ", " << weights << ", " << dst << ")");
vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_pre_f32;
const float scale = ggml_get_op_params_f32(dst, 0);
const bool gated = ggml_get_op_params_i32(dst, 1) != 0;
vk_pipeline pipeline = gated ? ctx->device->pipeline_dsv4_hc_pre_gated_f32 : ctx->device->pipeline_dsv4_hc_pre_f32;
GGML_ASSERT(pipeline != nullptr);
const uint32_t n_embd = (uint32_t)x->ne[0];
@@ -10371,9 +10380,10 @@ static void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subct
vk_op_dsv4_hc_pre_push_constants pc = {
n_embd, n_tokens,
ggml_vk_nb_elem(x, 0), ggml_vk_nb_elem(x, 1), ggml_vk_nb_elem(x, 2),
ggml_vk_nb_elem(weights, 0), ggml_vk_nb_elem(weights, 1),
ggml_vk_nb_elem(weights, 0), ggml_vk_nb_elem(weights, 1), ggml_vk_nb_elem(weights, 2),
ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1),
0, 0, 0,
scale,
};
init_pushconst_tensor_offsets(ctx, pc, x, weights, nullptr, nullptr, dst);
@@ -10383,7 +10393,7 @@ static void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subct
static void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst) {
VK_LOG_DEBUG("ggml_vk_dsv4_hc_post(" << x << ", " << residual << ", " << post << ", " << comb << ", " << dst << ")");
vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_post_f32;
vk_pipeline pipeline = comb ? ctx->device->pipeline_dsv4_hc_post_f32 : ctx->device->pipeline_dsv4_hc_post_nocomb_f32;
GGML_ASSERT(pipeline != nullptr);
const uint32_t n_embd = (uint32_t)x->ne[0];
@@ -10394,7 +10404,7 @@ static void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subc
const vk_subbuffer x_buf = ggml_vk_tensor_subbuffer(ctx, x, true);
const vk_subbuffer r_buf = ggml_vk_tensor_subbuffer(ctx, residual, true);
const vk_subbuffer p_buf = ggml_vk_tensor_subbuffer(ctx, post, true);
const vk_subbuffer c_buf = ggml_vk_tensor_subbuffer(ctx, comb, true);
const vk_subbuffer c_buf = comb ? ggml_vk_tensor_subbuffer(ctx, comb, true) : x_buf;
const vk_subbuffer d_buf = ggml_vk_tensor_subbuffer(ctx, dst, true);
vk_op_dsv4_hc_post_push_constants pc = {
@@ -10402,7 +10412,7 @@ static void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subc
ggml_vk_nb_elem(x, 0), ggml_vk_nb_elem(x, 1),
ggml_vk_nb_elem(residual, 0), ggml_vk_nb_elem(residual, 1), ggml_vk_nb_elem(residual, 2),
ggml_vk_nb_elem(post, 0), ggml_vk_nb_elem(post, 1),
ggml_vk_nb_elem(comb, 0), ggml_vk_nb_elem(comb, 1), ggml_vk_nb_elem(comb, 2),
comb ? ggml_vk_nb_elem(comb, 0) : 0, comb ? ggml_vk_nb_elem(comb, 1) : 0, comb ? ggml_vk_nb_elem(comb, 2) : 0,
ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1), ggml_vk_nb_elem(dst, 2),
0, 0, 0, 0, 0,
};
@@ -19692,10 +19702,10 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
}
// hc is hardcoded to 4 in the shaders. ggml only constrains it
// to 4 for COMB, so PRE/POST have to be checked here.
if (op->op == GGML_OP_DSV4_HC_PRE && (op->src[0]->ne[1] != 4 || ggml_get_op_params_i32(op, 1) != 0)) {
if (op->op == GGML_OP_DSV4_HC_PRE && op->src[0]->ne[1] != 4) {
return false;
}
if (op->op == GGML_OP_DSV4_HC_POST && (op->src[1]->ne[1] != 4 || op->src[3] == nullptr)) {
if (op->op == GGML_OP_DSV4_HC_POST && op->src[1]->ne[1] != 4) {
return false;
}
if (op->op == GGML_OP_DSV4_HC_COMB) {
@@ -20695,7 +20705,11 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph *
tensor_clone = ggml_dsv4_hc_comb(ggml_ctx, src_clone[0], src_clone[1], src_clone[2],
ggml_get_op_params_f32(tensor, 0), ggml_get_op_params_i32(tensor, 1));
} else if (tensor->op == GGML_OP_DSV4_HC_PRE) {
if (ggml_get_op_params_i32(tensor, 1) != 0) {
tensor_clone = ggml_dsv4_hc_pre_gated(ggml_ctx, src_clone[0], src_clone[1], ggml_get_op_params_f32(tensor, 0));
} else {
tensor_clone = ggml_dsv4_hc_pre(ggml_ctx, src_clone[0], src_clone[1]);
}
} else if (tensor->op == GGML_OP_DSV4_HC_POST) {
tensor_clone = ggml_dsv4_hc_post(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3]);
} else if (tensor->op == GGML_OP_MEAN) {
@@ -7,8 +7,13 @@
//
// dst[i0, idst, it] = x[i0, it]*post[idst, it]
// + sum_isrc residual[i0, isrc, it]*comb[idst, isrc, it]
//
// HAS_COMB == 0: identity mixing, each stream keeps its own residual:
//
// dst[i0, idst, it] = x[i0, it]*post[idst, it] + residual[i0, idst, it]
layout(constant_id = 0) const uint BLOCK_SIZE = 256;
layout(constant_id = 1) const uint HAS_COMB = 1;
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
@@ -48,7 +53,7 @@ void main() {
if (tid < hc) {
post_s[tid] = data_p[p_offset + tid * nbp0 + it * nbp1];
}
if (tid < hc * hc) {
if (HAS_COMB == 1 && tid < hc * hc) {
const uint idst = tid & 3;
const uint isrc = tid >> 2;
comb_s[tid] = data_c[c_offset + idst * nbc0 + isrc * nbc1 + it * nbc2];
@@ -74,10 +79,14 @@ void main() {
[[unroll]]
for (uint idst = 0; idst < hc; ++idst) {
float result = xv * post_s[idst];
if (HAS_COMB == 1) {
[[unroll]]
for (uint isrc = 0; isrc < hc; ++isrc) {
result = fma(r[isrc], comb_s[idst + hc * isrc], result);
}
} else {
result += r[idst];
}
data_d[d_offset + i0 * nbd0 + idst * nbd1 + it * nbd2] = result;
}
}
@@ -4,9 +4,14 @@
// Collapse the hc residual streams of a token into one, weighted per stream:
//
// dst[i0, it] = sum_ih x[i0, ih, it] * weights[ih, it]
// dst[i0, it] = scale * sum_ih x[i0, ih, it] * weights[ih, it]
//
// GATED: weights is a per-element gate [n_embd, hc, n_tokens], applied as sigmoid:
//
// dst[i0, it] = scale * sum_ih x[i0, ih, it] * sigmoid(gate[i0, ih, it])
layout(constant_id = 0) const uint BLOCK_SIZE = 256;
layout(constant_id = 1) const uint GATED = 0;
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
@@ -16,12 +21,14 @@ layout(push_constant) uniform parameter
uint n_tokens;
uint nbx0; uint nbx1; uint nbx2; // x
uint nbw0; uint nbw1; // weights
uint nbw0; uint nbw1; uint nbw2; // weights / gate
uint nbd0; uint nbd1; // dst
uint x_offset;
uint w_offset;
uint d_offset;
float scale;
};
layout(binding = 0, std430) readonly buffer X { float data_x[]; };
@@ -36,10 +43,12 @@ void main() {
const uint tid = gl_LocalInvocationID.x;
const uint it = gl_WorkGroupID.y;
if (GATED == 0) {
if (tid < hc) {
w[tid] = data_w[w_offset + tid * nbw0 + it * nbw1];
}
barrier();
}
// After the barrier, so every invocation reaches it.
const uint i0 = gl_WorkGroupID.x * BLOCK_SIZE + tid;
@@ -48,12 +57,19 @@ void main() {
}
const uint xb = x_offset + i0 * nbx0 + it * nbx2;
const uint wb = w_offset + i0 * nbw0 + it * nbw2;
float result = 0.0f;
[[unroll]]
for (uint ih = 0; ih < hc; ++ih) {
result = fma(data_x[xb + ih * nbx1], w[ih], result);
float wv;
if (GATED == 1) {
wv = 1.0f / (1.0f + exp(-data_w[wb + ih * nbw1]));
} else {
wv = w[ih];
}
result = fma(data_x[xb + ih * nbx1], wv, result);
}
data_d[d_offset + i0 * nbd0 + it * nbd1] = result;
data_d[d_offset + i0 * nbd0 + it * nbd1] = scale * result;
}