qwen4exp: add hc ops (llama/28901)

This commit is contained in:
Aman Gupta
2026-09-23 20:46:47 +03:00
committed by Georgi Gerganov
parent f9ad986899
commit fa2c801c5b
8 changed files with 133 additions and 43 deletions
+10
View File
@@ -2703,11 +2703,21 @@ extern "C" {
struct ggml_tensor * x,
struct ggml_tensor * weights);
// hc_pre with a per-element gate (Qwen3.8-Flash-Next): gate [n_embd, hc, n_tokens]
// result[i, t] = scale*sum_h x[i, h, t]*sigmoid(gate[i, h, t])
//
GGML_API struct ggml_tensor * ggml_dsv4_hc_pre_gated(
struct ggml_context * ctx,
struct ggml_tensor * x,
struct ggml_tensor * gate,
float scale);
// hc_post: x [n_embd, n_tokens], residual [n_embd, hc, n_tokens],
// post [hc, n_tokens], comb [dst_hc, src_hc, n_tokens]
// -> [n_embd, hc, n_tokens]
// result[i, dst, t] = x[i, t]*post[dst, t]
// + sum_src residual[i, src, t]*comb[dst, src, t]
// comb == NULL uses the identity: result[i, dst, t] = x[i, t]*post[dst, t] + residual[i, dst, t]
//
GGML_API struct ggml_tensor * ggml_dsv4_hc_post(
struct ggml_context * ctx,
+42 -14
View File
@@ -11259,10 +11259,19 @@ static void ggml_compute_forward_dsv4_hc_pre_f32(
const int64_t hc = x->ne[1];
const int64_t n_tokens = x->ne[2];
const float scale = ggml_get_op_params_f32(dst, 0);
const bool gated = ggml_get_op_params_i32(dst, 1) != 0;
GGML_ASSERT(dst->ne[0] == n_embd);
GGML_ASSERT(dst->ne[1] == n_tokens);
GGML_ASSERT(weights->ne[0] == hc);
GGML_ASSERT(weights->ne[1] == n_tokens);
if (gated) {
GGML_ASSERT(weights->ne[0] == n_embd);
GGML_ASSERT(weights->ne[1] == hc);
GGML_ASSERT(weights->ne[2] == n_tokens);
} else {
GGML_ASSERT(weights->ne[0] == hc);
GGML_ASSERT(weights->ne[1] == n_tokens);
}
GGML_TENSOR_LOCALS(size_t, nbx, x, nb);
GGML_TENSOR_LOCALS(size_t, nbw, weights, nb);
@@ -11282,12 +11291,18 @@ static void ggml_compute_forward_dsv4_hc_pre_f32(
float sum = 0.0f;
for (int64_t ih = 0; ih < hc; ++ih) {
const float xv = *(const float *) ((const char *) x->data + i0*nbx0 + ih*nbx1 + it*nbx2);
const float wv = *(const float *) ((const char *) weights->data + ih*nbw0 + it*nbw1);
const float xv = *(const float *) ((const char *) x->data + i0*nbx0 + ih*nbx1 + it*nbx2);
float wv;
if (gated) {
const float gv = *(const float *) ((const char *) weights->data + i0*nbw0 + ih*nbw1 + it*nbw2);
wv = 1.0f / (1.0f + expf(-gv));
} else {
wv = *(const float *) ((const char *) weights->data + ih*nbw0 + it*nbw1);
}
sum += xv * wv;
}
*(float *) ((char *) dst->data + i0*nbd0 + it*nbd1) = sum;
*(float *) ((char *) dst->data + i0*nbd0 + it*nbd1) = scale * sum;
}
}
@@ -11321,7 +11336,6 @@ static void ggml_compute_forward_dsv4_hc_post_f32(
GGML_ASSERT(x->type == GGML_TYPE_F32);
GGML_ASSERT(residual->type == GGML_TYPE_F32);
GGML_ASSERT(post->type == GGML_TYPE_F32);
GGML_ASSERT(comb->type == GGML_TYPE_F32);
GGML_ASSERT(dst->type == GGML_TYPE_F32);
const int64_t n_embd = x->ne[0];
@@ -11335,14 +11349,24 @@ static void ggml_compute_forward_dsv4_hc_post_f32(
GGML_ASSERT(residual->ne[2] == n_tokens);
GGML_ASSERT(post->ne[0] == hc);
GGML_ASSERT(post->ne[1] == n_tokens);
GGML_ASSERT(comb->ne[0] == hc);
GGML_ASSERT(comb->ne[1] == hc);
GGML_ASSERT(comb->ne[2] == n_tokens);
// comb == NULL: identity mixing, each stream keeps its own residual
size_t nbc0 = 0;
size_t nbc1 = 0;
size_t nbc2 = 0;
if (comb) {
GGML_ASSERT(comb->type == GGML_TYPE_F32);
GGML_ASSERT(comb->ne[0] == hc);
GGML_ASSERT(comb->ne[1] == hc);
GGML_ASSERT(comb->ne[2] == n_tokens);
nbc0 = comb->nb[0];
nbc1 = comb->nb[1];
nbc2 = comb->nb[2];
}
GGML_TENSOR_LOCALS(size_t, nbx, x, nb);
GGML_TENSOR_LOCALS(size_t, nbr, residual, nb);
GGML_TENSOR_LOCALS(size_t, nbp, post, nb);
GGML_TENSOR_LOCALS(size_t, nbc, comb, nb);
GGML_TENSOR_LOCALS(size_t, nbd, dst, nb);
const int ith = params->ith;
@@ -11362,10 +11386,14 @@ static void ggml_compute_forward_dsv4_hc_post_f32(
const float pv = *(const float *) ((const char *) post->data + idst*nbp0 + it*nbp1);
float sum = xv * pv;
for (int64_t isrc = 0; isrc < hc; ++isrc) {
const float rv = *(const float *) ((const char *) residual->data + i0*nbr0 + isrc*nbr1 + it*nbr2);
const float cv = *(const float *) ((const char *) comb->data + idst*nbc0 + isrc*nbc1 + it*nbc2);
sum += rv * cv;
if (comb) {
for (int64_t isrc = 0; isrc < hc; ++isrc) {
const float rv = *(const float *) ((const char *) residual->data + i0*nbr0 + isrc*nbr1 + it*nbr2);
const float cv = *(const float *) ((const char *) comb->data + idst*nbc0 + isrc*nbc1 + it*nbc2);
sum += rv * cv;
}
} else {
sum += *(const float *) ((const char *) residual->data + i0*nbr0 + idst*nbr1 + it*nbr2);
}
*(float *) ((char *) dst->data + i0*nbd0 + idst*nbd1 + it*nbd2) = sum;
+36 -14
View File
@@ -100,6 +100,7 @@ static __global__ void dsv4_hc_comb_f32(
}
}
template <bool gated>
static __global__ void dsv4_hc_pre_f32(
const float * x,
const float * weights,
@@ -112,8 +113,10 @@ static __global__ void dsv4_hc_pre_f32(
int64_t sx2,
int64_t sw0,
int64_t sw1,
int64_t sw2,
int64_t sd0,
int64_t sd1) {
int64_t sd1,
float scale) {
ggml_cuda_pdl_lc();
const int64_t ir = (int64_t) blockIdx.x * blockDim.x + threadIdx.x;
const int64_t nr = n_embd * n_tokens;
@@ -127,16 +130,22 @@ static __global__ void dsv4_hc_pre_f32(
const int64_t i0 = ir % n_embd;
const int64_t it = ir / n_embd;
float sum = x[i0*sx0 + it*sx2] * weights[it*sw1];
for (int64_t ih = 1; ih < hc; ++ih) {
float sum = 0.0f;
for (int64_t ih = 0; ih < hc; ++ih) {
const float xv = x[i0*sx0 + ih*sx1 + it*sx2];
const float wv = weights[ih*sw0 + it*sw1];
float wv;
if constexpr (gated) {
wv = 1.0f / (1.0f + expf(-weights[i0*sw0 + ih*sw1 + it*sw2]));
} else {
wv = weights[ih*sw0 + it*sw1];
}
sum += xv * wv;
}
dst[i0*sd0 + it*sd1] = sum;
dst[i0*sd0 + it*sd1] = scale * sum;
}
template <bool has_comb>
static __global__ void dsv4_hc_post_f32(
const float * x,
const float * residual,
@@ -174,8 +183,12 @@ static __global__ void dsv4_hc_post_f32(
const int64_t it = ir / (n_embd * hc);
float sum = x[i0*sx0 + it*sx1] * post[idst*sp0 + it*sp1];
for (int64_t isrc = 0; isrc < hc; ++isrc) {
sum += residual[i0*sr0 + isrc*sr1 + it*sr2] * comb[idst*sc0 + isrc*sc1 + it*sc2];
if constexpr (has_comb) {
for (int64_t isrc = 0; isrc < hc; ++isrc) {
sum += residual[i0*sr0 + isrc*sr1 + it*sr2] * comb[idst*sc0 + isrc*sc1 + it*sc2];
}
} else {
sum += residual[i0*sr0 + idst*sr1 + it*sr2];
}
dst[i0*sd0 + idst*sd1 + it*sd2] = sum;
@@ -240,18 +253,23 @@ void ggml_cuda_op_dsv4_hc_pre(ggml_backend_cuda_context & ctx, ggml_tensor * dst
const int64_t hc = x->ne[1];
const int64_t n_tokens = x->ne[2];
const float scale = ggml_get_op_params_f32(dst, 0);
const bool gated = ggml_get_op_params_i32(dst, 1) != 0;
const int block_size = 256;
const int64_t nr = n_embd * n_tokens;
const dim3 block_dims(block_size, 1, 1);
const dim3 grid_dims((nr + block_size - 1) / block_size, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream());
ggml_cuda_kernel_launch(dsv4_hc_pre_f32, launch_params,
auto kernel = gated ? dsv4_hc_pre_f32<true> : dsv4_hc_pre_f32<false>;
ggml_cuda_kernel_launch(kernel, launch_params,
(const float *) x->data, (const float *) weights->data, (float *) dst->data,
n_embd, hc, n_tokens,
nbx0 / sizeof(float), nbx1 / sizeof(float), nbx2 / sizeof(float),
nbw0 / sizeof(float), nbw1 / sizeof(float),
nbd0 / sizeof(float), nbd1 / sizeof(float));
nbw0 / sizeof(float), nbw1 / sizeof(float), nbw2 / sizeof(float),
nbd0 / sizeof(float), nbd1 / sizeof(float),
scale);
}
void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
@@ -263,15 +281,18 @@ void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * ds
GGML_ASSERT(x->type == GGML_TYPE_F32);
GGML_ASSERT(residual->type == GGML_TYPE_F32);
GGML_ASSERT(post->type == GGML_TYPE_F32);
GGML_ASSERT(comb->type == GGML_TYPE_F32);
GGML_ASSERT(comb == nullptr || comb->type == GGML_TYPE_F32);
GGML_ASSERT(dst->type == GGML_TYPE_F32);
GGML_TENSOR_LOCALS(size_t, nbx, x, nb);
GGML_TENSOR_LOCALS(size_t, nbr, residual, nb);
GGML_TENSOR_LOCALS(size_t, nbp, post, nb);
GGML_TENSOR_LOCALS(size_t, nbc, comb, nb);
GGML_TENSOR_LOCALS(size_t, nbd, dst, nb);
const size_t nbc0 = comb ? comb->nb[0] : 0;
const size_t nbc1 = comb ? comb->nb[1] : 0;
const size_t nbc2 = comb ? comb->nb[2] : 0;
const int64_t n_embd = x->ne[0];
const int64_t n_tokens = x->ne[1];
const int64_t hc = residual->ne[1];
@@ -282,9 +303,10 @@ void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * ds
const dim3 grid_dims((nr + block_size - 1) / block_size, 1, 1);
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream());
ggml_cuda_kernel_launch(dsv4_hc_post_f32, launch_params,
auto kernel = comb ? dsv4_hc_post_f32<true> : dsv4_hc_post_f32<false>;
ggml_cuda_kernel_launch(kernel, launch_params,
(const float *) x->data, (const float *) residual->data,
(const float *) post->data, (const float *) comb->data, (float *) dst->data,
(const float *) post->data, comb ? (const float *) comb->data : nullptr, (float *) dst->data,
n_embd, hc, n_tokens,
nbx0 / sizeof(float), nbx1 / sizeof(float),
nbr0 / sizeof(float), nbr1 / sizeof(float), nbr2 / sizeof(float),
+1 -1
View File
@@ -5497,7 +5497,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
op->type == GGML_TYPE_F32;
case GGML_OP_DSV4_HC_POST:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && (op->src[3] == nullptr || op->src[3]->type == GGML_TYPE_F32) &&
op->type == GGML_TYPE_F32;
case GGML_OP_FLASH_ATTN_EXT:
return ggml_cuda_flash_attn_ext_supported(dev_ctx->device, op);
+2
View File
@@ -1803,6 +1803,7 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
op->type == GGML_TYPE_F32 &&
op->src[0]->ne[1] == 4 &&
op->src[1]->ne[0] == 4 &&
op->src[1]->ne[2] == 1 &&
ggml_is_contiguous_rows(op->src[0]) &&
ggml_is_contiguous_rows(op->src[1]);
case GGML_OP_DSV4_HC_POST:
@@ -1810,6 +1811,7 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
op->src[0]->type == GGML_TYPE_F32 &&
op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 &&
op->src[3] != NULL &&
op->src[3]->type == GGML_TYPE_F32 &&
op->type == GGML_TYPE_F32 &&
op->src[1]->ne[1] == 4 &&
+2 -2
View File
@@ -6434,13 +6434,13 @@ static bool do_ggml_backend_sycl_device_supports_op(ggml_backend_dev_t dev, cons
break;
case GGML_OP_DSV4_HC_PRE:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->type == GGML_TYPE_F32;
op->type == GGML_TYPE_F32 && ggml_get_op_params_i32(op, 1) == 0;
case GGML_OP_DSV4_HC_COMB:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
case GGML_OP_DSV4_HC_POST:
return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 &&
op->src[2]->type == GGML_TYPE_F32 && op->src[3] != nullptr && op->src[3]->type == GGML_TYPE_F32 &&
op->type == GGML_TYPE_F32;
case GGML_OP_LIGHTNING_INDEXER:
return op->src[0]->type == GGML_TYPE_F32 &&
+2 -2
View File
@@ -19692,10 +19692,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) {
if (op->op == GGML_OP_DSV4_HC_PRE && (op->src[0]->ne[1] != 4 || ggml_get_op_params_i32(op, 1) != 0)) {
return false;
}
if (op->op == GGML_OP_DSV4_HC_POST && op->src[1]->ne[1] != 4) {
if (op->op == GGML_OP_DSV4_HC_POST && (op->src[1]->ne[1] != 4 || op->src[3] == nullptr)) {
return false;
}
if (op->op == GGML_OP_DSV4_HC_COMB) {
+38 -10
View File
@@ -6507,10 +6507,12 @@ struct ggml_tensor * ggml_dsv4_hc_comb(
// ggml_dsv4_hc_pre
struct ggml_tensor * ggml_dsv4_hc_pre(
static struct ggml_tensor * ggml_dsv4_hc_pre_impl(
struct ggml_context * ctx,
struct ggml_tensor * x,
struct ggml_tensor * weights) {
struct ggml_tensor * weights,
float scale,
bool gated) {
GGML_ASSERT(x->type == GGML_TYPE_F32);
GGML_ASSERT(weights->type == GGML_TYPE_F32);
@@ -6520,13 +6522,22 @@ struct ggml_tensor * ggml_dsv4_hc_pre(
GGML_ASSERT(hc > 0);
GGML_ASSERT(x->ne[3] == 1);
GGML_ASSERT(weights->ne[0] == hc);
GGML_ASSERT(weights->ne[1] == n_tokens);
GGML_ASSERT(weights->ne[2] == 1);
if (gated) {
GGML_ASSERT(weights->ne[0] == n_embd);
GGML_ASSERT(weights->ne[1] == hc);
GGML_ASSERT(weights->ne[2] == n_tokens);
} else {
GGML_ASSERT(weights->ne[0] == hc);
GGML_ASSERT(weights->ne[1] == n_tokens);
GGML_ASSERT(weights->ne[2] == 1);
}
GGML_ASSERT(weights->ne[3] == 1);
struct ggml_tensor * result = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens);
ggml_set_op_params_f32(result, 0, scale);
ggml_set_op_params_i32(result, 1, gated ? 1 : 0);
result->op = GGML_OP_DSV4_HC_PRE;
result->src[0] = x;
result->src[1] = weights;
@@ -6534,6 +6545,21 @@ struct ggml_tensor * ggml_dsv4_hc_pre(
return result;
}
struct ggml_tensor * ggml_dsv4_hc_pre(
struct ggml_context * ctx,
struct ggml_tensor * x,
struct ggml_tensor * weights) {
return ggml_dsv4_hc_pre_impl(ctx, x, weights, 1.0f, false);
}
struct ggml_tensor * ggml_dsv4_hc_pre_gated(
struct ggml_context * ctx,
struct ggml_tensor * x,
struct ggml_tensor * gate,
float scale) {
return ggml_dsv4_hc_pre_impl(ctx, x, gate, scale, true);
}
// ggml_dsv4_hc_post
struct ggml_tensor * ggml_dsv4_hc_post(
@@ -6545,7 +6571,6 @@ struct ggml_tensor * ggml_dsv4_hc_post(
GGML_ASSERT(x->type == GGML_TYPE_F32);
GGML_ASSERT(residual->type == GGML_TYPE_F32);
GGML_ASSERT(post->type == GGML_TYPE_F32);
GGML_ASSERT(comb->type == GGML_TYPE_F32);
const int64_t n_embd = x->ne[0];
const int64_t n_tokens = x->ne[1];
@@ -6564,10 +6589,13 @@ struct ggml_tensor * ggml_dsv4_hc_post(
GGML_ASSERT(post->ne[2] == 1);
GGML_ASSERT(post->ne[3] == 1);
GGML_ASSERT(comb->ne[0] == hc);
GGML_ASSERT(comb->ne[1] == hc);
GGML_ASSERT(comb->ne[2] == n_tokens);
GGML_ASSERT(comb->ne[3] == 1);
if (comb) {
GGML_ASSERT(comb->type == GGML_TYPE_F32);
GGML_ASSERT(comb->ne[0] == hc);
GGML_ASSERT(comb->ne[1] == hc);
GGML_ASSERT(comb->ne[2] == n_tokens);
GGML_ASSERT(comb->ne[3] == 1);
}
struct ggml_tensor * result = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens);