diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 45e9537f0e45..73c5557d17ca 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -4500,24 +4500,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(match.experts), match.dst); + params->add_alloc_dep(params->user_data, const_cast(match.weights), match.dst); + if (match.expert_scale != nullptr) { + params->add_alloc_dep( + params->user_data, const_cast(match.expert_scale), match.dst); + } + i += match.node_count - 1; } - params->add_alloc_dep(params->user_data, const_cast(match.experts), match.dst); - params->add_alloc_dep(params->user_data, const_cast(match.weights), match.dst); - if (match.expert_scale != nullptr) { - params->add_alloc_dep( - params->user_data, const_cast(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 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; } }