diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index 85a1ae7ae..224bdef92 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -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, diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 266261c5e..23001254c 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -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; diff --git a/ggml/src/ggml-cuda/dsv4-hc.cu b/ggml/src/ggml-cuda/dsv4-hc.cu index c4b19a787..ca1d2dc8a 100644 --- a/ggml/src/ggml-cuda/dsv4-hc.cu +++ b/ggml/src/ggml-cuda/dsv4-hc.cu @@ -100,6 +100,7 @@ static __global__ void dsv4_hc_comb_f32( } } +template 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 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 : dsv4_hc_pre_f32; + 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 : dsv4_hc_post_f32; + 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), diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 74bb47145..a9038f1f4 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -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); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 5654c5004..c734c8e13 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -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 && diff --git a/ggml/src/ggml-sycl/ggml-sycl.cpp b/ggml/src/ggml-sycl/ggml-sycl.cpp index 686a4c76e..6f9ead60e 100644 --- a/ggml/src/ggml-sycl/ggml-sycl.cpp +++ b/ggml/src/ggml-sycl/ggml-sycl.cpp @@ -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 && diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index f936127a6..baa44ad1f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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) { diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 5ef03e190..175281409 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -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);