Commit 1a679828f for llama.cpp
commit 1a679828f3312ebc53c9285805cc8bd7c7d90366
Author: Aman Gupta <amangupta052@gmail.com>
Date: Wed Sep 23 13:26:05 2026 +0800
cuda: top-k MoE should always fire (#28432)
diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu
index c8b23b2d5..60d72046e 100644
--- a/ggml/src/ggml-cuda/ggml-cuda.cu
+++ b/ggml/src/ggml-cuda/ggml-cuda.cu
@@ -4505,24 +4505,90 @@ static void ggml_backend_cuda_graph_optimize(ggml_backend_t backend, ggml_cgraph
ggml_backend_cuda_context * cuda_ctx = (ggml_backend_cuda_context *) backend->context;
static const bool disable_fusion = getenv("GGML_CUDA_DISABLE_FUSION") != nullptr && std::atoi(getenv("GGML_CUDA_DISABLE_FUSION"));
- if (!disable_fusion) {
- for (int i = 0; i < cgraph->n_nodes; ++i) {
- if (cgraph->nodes[i]->op != GGML_OP_MUL) {
- continue;
+
+ auto add_alloc_deps = [&](size_t start, size_t last_node) {
+
+ for (size_t i = start; i < last_node; ++i) {
+ params->add_alloc_dep(params->user_data, cgraph->nodes[i], cgraph->nodes[last_node]);
+
+ for (int j = 0; j < GGML_MAX_SRC; ++j) {
+ if (cgraph->nodes[i]->src[j]) {
+ params->add_alloc_dep(params->user_data, cgraph->nodes[i]->src[j], cgraph->nodes[last_node]);
+ }
}
+ }
+ };
+ if (!disable_fusion) {
+ // add alloc deps for performance positive fusions. This may increase the overall compute buffer size.
+ // TODO: consolidate fusion paths in graph_optimize and graph_compute
+ for (int i = 0; i < cgraph->n_nodes; ++i) {
ggml_cuda_moe_weighted_reduction_match match;
- if (!ggml_cuda_match_moe_weighted_reduction(cgraph, i, match)) {
- continue;
+ if (ggml_cuda_match_moe_weighted_reduction(cgraph, i, match)) {
+ params->add_alloc_dep(params->user_data, const_cast<ggml_tensor *>(match.experts), match.dst);
+ params->add_alloc_dep(params->user_data, const_cast<ggml_tensor *>(match.weights), match.dst);
+ if (match.expert_scale != nullptr) {
+ params->add_alloc_dep(
+ params->user_data, const_cast<ggml_tensor *>(match.expert_scale), match.dst);
+ }
+ i += match.node_count - 1;
}
- params->add_alloc_dep(params->user_data, const_cast<ggml_tensor *>(match.experts), match.dst);
- params->add_alloc_dep(params->user_data, const_cast<ggml_tensor *>(match.weights), match.dst);
- if (match.expert_scale != nullptr) {
- params->add_alloc_dep(
- params->user_data, const_cast<ggml_tensor *>(match.expert_scale), match.dst);
+ if (cgraph->nodes[i]->op == GGML_OP_UNARY || cgraph->nodes[i]->op == GGML_OP_SOFT_MAX ||
+ cgraph->nodes[i]->op == GGML_OP_ARGSORT) {
+ ggml_cuda_topk_moe_args args;
+ const bool can_fuse = ggml_cuda_topk_moe_fusion(cgraph, i, args);
+ std::vector<ggml_op> ops;
+
+ const ggml_tensor * node = cgraph->nodes[i];
+
+ if (can_fuse) {
+ const ggml_tensor * logits = node->src[0];
+ ggml_tensor * weights = nullptr;
+ ggml_tensor * ids = nullptr;
+
+ if (!args.delayed_softmax) {
+ int out_nodes[2]; // nodes which can't be elided
+
+ if (args.sigmoid) {
+ ops.insert(ops.end(), { GGML_OP_UNARY });
+ } else if (args.sqrt_softplus) {
+ ops.insert(ops.end(), { GGML_OP_UNARY, GGML_OP_SQRT });
+ } else {
+ ops.insert(ops.end(), { GGML_OP_SOFT_MAX });
+ }
+ const int i_probs = i + (int) ops.size() - 1; // last node of the gating activation
+
+ if (args.prob_bias) {
+ ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ADD, GGML_OP_ARGSORT, GGML_OP_VIEW,
+ GGML_OP_GET_ROWS });
+ out_nodes[0] = i_probs + 4;
+ } else {
+ ops.insert(ops.end(), { GGML_OP_RESHAPE, GGML_OP_ARGSORT, GGML_OP_VIEW, GGML_OP_GET_ROWS });
+ out_nodes[0] = i_probs + 3;
+ }
+ ids = cgraph->nodes[out_nodes[0]];
+
+ if (args.norm) {
+ ops.insert(ops.end(),
+ { GGML_OP_RESHAPE, GGML_OP_SUM_ROWS, GGML_OP_CLAMP, GGML_OP_DIV, GGML_OP_RESHAPE });
+ }
+ if (args.scale) {
+ ops.insert(ops.end(), { GGML_OP_SCALE });
+ }
+
+ weights = cgraph->nodes[i + ops.size() - 1];
+ out_nodes[1] = i + ops.size() - 1;
+
+ if (ggml_can_fuse_subgraph(cgraph, i, ops.size(), ops.data(), out_nodes, 2) &&
+ ggml_cuda_should_use_topk_moe(node, logits, weights, ids)) {
+
+ add_alloc_deps(i, i + ops.size());
+ i += ops.size() - 1;
+ }
+ }
+ }
}
- i += match.node_count - 1;
}
}