perf: reduce CPU overhead in graph execution and sampling (#1997)

This commit is contained in:
leejet
2026-09-19 18:32:58 +08:00
committed by GitHub
parent d32b4e893b
commit 275ab58e01
8 changed files with 195 additions and 123 deletions
+24 -27
View File
@@ -342,21 +342,17 @@ void GGMLRunner::copy_data_to_backend_tensor(ggml_cgraph* gf, bool clear_after_c
}
}
bool GGMLRunner::resolve_graph_cut_plan(ggml_cgraph* gf,
GraphCutPlan* plan_out) {
GGML_ASSERT(plan_out != nullptr);
const GGMLRunner::GraphCutPlan& GGMLRunner::resolve_graph_cut_plan(ggml_cgraph* gf) {
GGML_ASSERT(gf != nullptr);
*plan_out = sd::ggml_graph_cut::resolve_plan(runtime_backend,
gf,
&graph_cut_plan_cache_,
params_tensor_set_,
get_desc().c_str());
return true;
return sd::ggml_graph_cut::resolve_plan(runtime_backend,
gf,
&graph_cut_plan_cache_,
params_tensor_set_,
get_desc().c_str());
}
bool GGMLRunner::resolve_graph_cut_layer_split_plan(ggml_cgraph* gf,
GraphCutPlan* plan_out) {
return resolve_graph_cut_plan(gf, plan_out);
const GGMLRunner::GraphCutPlan& GGMLRunner::resolve_graph_cut_layer_split_plan(ggml_cgraph* gf) {
return resolve_graph_cut_plan(gf);
}
bool GGMLRunner::assign_graph_cut_layer_split_backends(ggml_cgraph* gf) {
@@ -369,10 +365,7 @@ bool GGMLRunner::assign_graph_cut_layer_split_backends(ggml_cgraph* gf) {
return false;
}
GraphCutPlan plan;
if (!resolve_graph_cut_layer_split_plan(gf, &plan)) {
return false;
}
const auto& plan = resolve_graph_cut_layer_split_plan(gf);
if (!plan.valid || !plan.has_cuts || plan.segments.size() <= 1) {
auto manager = residency_manager.lock();
if (manager == nullptr) {
@@ -820,24 +813,25 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
if (!assign_graph_cut_layer_split_backends(graph)) {
return std::nullopt;
}
const auto params = collect_used_param_tensors(graph);
ggml_graph_cut::Plan plan;
if (!resolve_graph_cut_plan(graph, &plan)) {
return std::nullopt;
}
const auto full_measurement = measure(graph, plan.compute_buffer_size);
const auto params = collect_used_param_tensors(graph);
const auto& cached_plan = resolve_graph_cut_plan(graph);
const auto full_measurement = measure(graph, cached_plan.compute_buffer_size);
if (full_measurement.buffers.empty()) {
return std::nullopt;
}
auto manager = residency_manager.lock();
const bool segmented = !is_multi_device() && !sd_backend_is_cpu(runtime_backend) &&
manager != nullptr && manager->segmented_compute_enabled() &&
plan.valid && plan.has_cuts && plan.segments.size() > 1 &&
cached_plan.valid && cached_plan.has_cuts && cached_plan.segments.size() > 1 &&
!fits(memory_requests(full_measurement.buffers, cache_.pending_bytes(graph)), params);
ggml_graph_cut::Plan monolithic_plan;
if (!segmented) {
ggml_graph_cut::Segment segment;
monolithic_plan.segments.emplace_back();
auto& segment = monolithic_plan.segments.back();
segment.group_name = "graph";
segment.compute_buffer_size = plan.compute_buffer_size;
segment.compute_buffer_size = cached_plan.compute_buffer_size;
segment.internal_node_indices.reserve(ggml_graph_n_nodes(graph));
segment.input_refs.reserve(ggml_graph_cut::leaf_count(graph));
for (int i = 0; i < ggml_graph_n_nodes(graph); ++i) {
segment.internal_node_indices.push_back(i);
}
@@ -850,8 +844,8 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
: ggml_graph_cut::Segment::INPUT_EXTERNAL;
segment.input_refs.push_back(input);
}
plan.segments = {std::move(segment)};
}
const auto& plan = segmented ? cached_plan : monolithic_plan;
const bool segments_changed = plan.segments.size() != logged_segment_count_;
if (segments_changed && (segmented || logged_segment_count_ > 1)) {
LOG_VERBOSE("%s using %zu segment%s", get_desc().c_str(),
@@ -910,7 +904,10 @@ std::optional<Tensor<float>> GGMLRunner::execute_graph(ggml_cgraph* graph, int n
auto ensure_capacity = [&]() {
sync_runtime_residency();
auto requests = memory_requests(measurement.buffers, new_cache_bytes);
if (!fits(requests, weights.params(index)) && workspace_.release_excess(measurement)) {
if (fits(requests, weights.params(index))) {
return true;
}
if (workspace_.release_excess(measurement)) {
sync_runtime_residency();
requests = memory_requests(measurement.buffers, new_cache_bytes);
}