struct ggml_cgraph * cgraph,
struct ggml_tensor * tensor);
+ // add the tensor and its parents to the graph without marking them for compute
+ // the flag is set later, when the tensor is reached from a node that computes
+ GGML_API void ggml_build_forward_order(
+ struct ggml_cgraph * cgraph,
+ struct ggml_tensor * tensor);
+
GGML_API void ggml_build_backward_expand(
struct ggml_context * ctx, // context for gradient computation
struct ggml_cgraph * cgraph,
ggml_build_forward_impl(cgraph, tensor, true, true);
}
+void ggml_build_forward_order(struct ggml_cgraph * cgraph, struct ggml_tensor * tensor) {
+ ggml_build_forward_impl(cgraph, tensor, true, false);
+}
+
void ggml_build_backward_expand(
struct ggml_context * ctx,
struct ggml_cgraph * cgraph,
ggml_tensor * sinks) const {
// these nodes are added to the graph together so that they are not reordered
// by doing so, the number of splits in the graph is reduced
- ggml_build_forward_expand(gf, q_cur);
- ggml_build_forward_expand(gf, k_cur);
- ggml_build_forward_expand(gf, v_cur);
+ // the order is fixed without the compute flag, so an unselected branch stays out of the compute set
+ ggml_build_forward_order(gf, q_cur);
+ ggml_build_forward_order(gf, k_cur);
+ ggml_build_forward_order(gf, v_cur);
ggml_tensor * q = ggml_permute(ctx0, q_cur, 0, 2, 1, 3);
//cb(q, "q", il);