Commit 1b0ba10 for stable-diffusion.cpp
commit 1b0ba10893f4e2e0656103011fc1c4645a02a734
Author: Daniele <57776841+daniandtheweb@users.noreply.github.com>
Date: Sat Oct 10 17:35:58 2026 +0200
feat: run conditional and unconditional CFG in one batched UNet forward (#2085)
diff --git a/docs/performance.md b/docs/performance.md
index 6797ad4..acd071c 100644
--- a/docs/performance.md
+++ b/docs/performance.md
@@ -21,6 +21,16 @@ CPU fallback. It excludes weights and cache buffers. Within a runner lifecycle,
the summary is printed only on the first graph or when backend capacities or the
segment count change.
+## Run conditional and unconditional CFG in one batched UNet forward.
+
+For UNet models, the conditional and unconditional guidance branches are
+concatenated into a single batch of two and run through one UNet forward per
+step instead of two separate forwards. This is enabled by default whenever the
+run qualifies for it.
+
+Use `--batched-cfg off` to force separate conditional and unconditional
+forwards.
+
## Use VAE tiling to reduce encode and decode memory usage.
`--vae-tiling` enables spatial tiling for both VAE encoding and decoding. The
diff --git a/examples/common/common.cpp b/examples/common/common.cpp
index 89b226f..faa6372 100644
--- a/examples/common/common.cpp
+++ b/examples/common/common.cpp
@@ -378,6 +378,23 @@ static int parse_scale_override(int argc, const char** argv, int index, float& s
return 1;
}
+static int parse_on_off_arg(int argc, const char** argv, int index, const char* option, bool& value) {
+ if (++index >= argc) {
+ LOG_ERROR("%s requires 'on' or 'off'", option);
+ return -1;
+ }
+ const std::string arg = argv[index];
+ if (arg == "on") {
+ value = true;
+ } else if (arg == "off") {
+ value = false;
+ } else {
+ LOG_ERROR("invalid %s value '%s'; expected 'on' or 'off'", option, argv[index]);
+ return -1;
+ }
+ return 1;
+}
+
ArgOptions SDContextParams::get_options() {
ArgOptions options;
options.string_options = {
@@ -637,23 +654,6 @@ ArgOptions SDContextParams::get_options() {
true, &vae_conv_direct},
};
- auto on_auto_fit_arg = [&](int argc, const char** argv, int index) {
- if (++index >= argc) {
- LOG_ERROR("--auto-fit requires 'on' or 'off'");
- return -1;
- }
- const std::string arg = argv[index];
- if (arg == "on") {
- auto_fit = true;
- } else if (arg == "off") {
- auto_fit = false;
- } else {
- LOG_ERROR("invalid --auto-fit value '%s'; expected 'on' or 'off'", argv[index]);
- return -1;
- }
- return 1;
- };
-
auto on_type_arg = [&](int argc, const char** argv, int index) {
if (++index >= argc) {
return -1;
@@ -742,7 +742,15 @@ ArgOptions SDContextParams::get_options() {
"on|off (default: on). Preserve --backend (otherwise select one GPU) and place weights on the compute GPU, "
"RAM, another GPU, or disk in that order, according to available memory (--max-vram limits GPU budgets). "
"Disabled by explicit --params-backend; uses automatic graph segmentation when needed",
- on_auto_fit_arg},
+ [this](int argc, const char** argv, int index) {
+ return parse_on_off_arg(argc, argv, index, "--auto-fit", auto_fit);
+ }},
+ {"",
+ "--batched-cfg",
+ "on|off (default: on). Run the conditional and unconditional CFG branches in one batched UNet forward when supported",
+ [this](int argc, const char** argv, int index) {
+ return parse_on_off_arg(argc, argv, index, "--batched-cfg", batched_cfg);
+ }},
{"",
"--type",
"weight type (examples: f32, f16, q4_0, q4_1, q5_0, q5_1, q8_0, q2_K, q3_K, q4_K). "
@@ -940,6 +948,7 @@ std::string SDContextParams::to_string() const {
<< " max_vram: \"" << max_vram << "\",\n"
<< " disable_prefetch: " << (disable_prefetch ? "true" : "false") << ",\n"
<< " disable_segmented_compute: " << (disable_segmented_compute ? "true" : "false") << ",\n"
+ << " batched_cfg: " << (batched_cfg ? "true" : "false") << ",\n"
<< " eager_load: " << (eager_load ? "true" : "false") << ",\n"
<< " backend: \"" << backend << "\",\n"
<< " params_backend: \"" << params_backend << "\",\n"
@@ -1022,6 +1031,7 @@ sd_ctx_params_t SDContextParams::to_sd_ctx_params_t(bool taesd_preview) {
sd_ctx_params.max_vram = max_vram.c_str();
sd_ctx_params.disable_prefetch = disable_prefetch;
sd_ctx_params.disable_segmented_compute = disable_segmented_compute;
+ sd_ctx_params.batched_cfg = batched_cfg;
sd_ctx_params.eager_load = eager_load;
sd_ctx_params.backend = effective_backend.c_str();
sd_ctx_params.params_backend = effective_params_backend.c_str();
diff --git a/examples/common/common.h b/examples/common/common.h
index 8b60c76..bd9075c 100644
--- a/examples/common/common.h
+++ b/examples/common/common.h
@@ -157,6 +157,7 @@ struct SDContextParams {
bool disable_prefetch = false;
bool disable_segmented_compute = false;
bool eager_load = false;
+ bool batched_cfg = true;
std::string backend;
std::string params_backend;
std::string split_mode;
diff --git a/include/stable-diffusion.h b/include/stable-diffusion.h
index a2e1575..c21f765 100644
--- a/include/stable-diffusion.h
+++ b/include/stable-diffusion.h
@@ -245,6 +245,7 @@ typedef struct {
const char* rpc_servers;
const char* model_args;
bool disable_segmented_compute; // Force monolithic graph execution even when automatic graph cutting would fit memory better
+ bool batched_cfg; // Run the conditional and unconditional CFG branches in one batched UNet forward when supported
float linear_scale; // Override linear input scaling; 0 keeps the model default
float attn_scale; // Override flash-attention K/V scaling; 0 keeps the model default
const char* tokenizer; // tokenizer.json path or main=FILE,clip-l=FILE,clip-g=FILE assignments; required for PiD and Lens
diff --git a/src/model/diffusion/unet.hpp b/src/model/diffusion/unet.hpp
index 1c8ddf4..7c0ff09 100644
--- a/src/model/diffusion/unet.hpp
+++ b/src/model/diffusion/unet.hpp
@@ -587,7 +587,9 @@ public:
label_emb = ggml_silu_inplace(ctx->ggml_ctx, label_emb);
label_emb = label_embed_2->forward(ctx, label_emb); // [N, time_embed_dim]
- emb = ggml_add(ctx->ggml_ctx, emb, label_emb); // [N, time_embed_dim]
+ emb = label_emb->ne[1] > emb->ne[1]
+ ? ggml_add(ctx->ggml_ctx, label_emb, emb)
+ : ggml_add(ctx->ggml_ctx, emb, label_emb); // [N, time_embed_dim]
}
// sd::ggml_graph_cut::mark_graph_cut(emb, "unet.prelude", "emb");
diff --git a/src/pipeline/diffusion_engine.cpp b/src/pipeline/diffusion_engine.cpp
index 68129bb..6400b24 100644
--- a/src/pipeline/diffusion_engine.cpp
+++ b/src/pipeline/diffusion_engine.cpp
@@ -2227,6 +2227,27 @@ void StableDiffusionGGML::report_sample_progress(int step,
}
}
+static sd::Tensor<float> batch_two_condition_tensors(const sd::Tensor<float>& a, const sd::Tensor<float>& b) {
+ if (a.empty() || b.empty() || a.dim() != b.dim()) {
+ return {};
+ }
+ if (a.dim() == 1) {
+ if (a.shape() != b.shape()) {
+ return {};
+ }
+ auto batched = sd::ops::concat(a, b, 0);
+ batched.reshape_({a.shape()[0], 2});
+ return batched;
+ }
+ const int64_t batch_dim = a.dim() - 1;
+ for (int64_t d = 0; d < batch_dim; d++) {
+ if (a.shape()[d] != b.shape()[d]) {
+ return {};
+ }
+ }
+ return sd::ops::concat(a, b, static_cast<size_t>(batch_dim));
+}
+
void StableDiffusionGGML::compute_sample_controls(const sd::Tensor<float>& control_image,
const sd::Tensor<float>& noised_input,
const sd::Tensor<float>& timesteps_tensor,
@@ -2604,6 +2625,57 @@ sd::Tensor<float> StableDiffusionGGML::sample(const std::shared_ptr<DiffusionMod
return output_opt;
};
+ auto run_batched_condition = [&](const SDCondition& condition,
+ const sd::Tensor<float>* c_concat_override) -> sd::Tensor<float> {
+ const sd::Tensor<float>& condition_concat =
+ c_concat_override != nullptr ? *c_concat_override : condition.c_concat;
+
+ sd::Tensor<float> batched_context = batch_two_condition_tensors(condition.c_crossattn, uncond.c_crossattn);
+ sd::Tensor<float> batched_y = batch_two_condition_tensors(condition.c_vector, uncond.c_vector);
+ sd::Tensor<float> batched_concat = batch_two_condition_tensors(condition_concat, uncond.c_concat);
+ if (!condition.c_crossattn.empty() && batched_context.empty()) {
+ return {};
+ }
+ if ((!condition.c_vector.empty() || !uncond.c_vector.empty()) && batched_y.empty()) {
+ return {};
+ }
+ if ((!condition_concat.empty() || !uncond.c_concat.empty()) && batched_concat.empty()) {
+ return {};
+ }
+
+ std::vector<sd::Tensor<float>> uncond_controls;
+ compute_sample_controls(control_image, noised_input, timesteps_tensor, uncond, &uncond_controls);
+ if (controls.size() != uncond_controls.size()) {
+ return {};
+ }
+ std::vector<sd::Tensor<float>> batched_controls;
+ batched_controls.reserve(controls.size());
+ for (size_t i = 0; i < controls.size(); i++) {
+ sd::Tensor<float> batched_control = batch_two_condition_tensors(controls[i], uncond_controls[i]);
+ if (batched_control.empty()) {
+ return {};
+ }
+ batched_controls.push_back(std::move(batched_control));
+ }
+
+ sd::Tensor<float> batched_x =
+ sd::ops::concat(noised_input, noised_input, static_cast<size_t>(noised_input.dim() - 1));
+
+ DiffusionParams batched_params = diffusion_params;
+ batched_params.x = &batched_x;
+ batched_params.context = batched_context.empty() ? nullptr : &batched_context;
+ batched_params.c_concat = batched_concat.empty() ? nullptr : &batched_concat;
+ batched_params.y = batched_y.empty() ? nullptr : &batched_y;
+ batched_params.ref_latents = nullptr;
+ batched_params.extra = UNetDiffusionExtra{1, &batched_controls, control_strength};
+
+ sd::Tensor<float> output = work_diffusion_model->compute(n_threads, batched_params);
+ if (output.empty()) {
+ LOG_ERROR("batched diffusion model compute failed");
+ }
+ return output;
+ };
+
const SDCondition* positive_condition = &cond;
const sd::Tensor<float>* c_concat_override = nullptr;
for (const auto& extension : generation_extensions) {
@@ -2643,12 +2715,40 @@ sd::Tensor<float> StableDiffusionGGML::sample(const std::shared_ptr<DiffusionMod
}
}
- cond_out = run_condition(*positive_condition, c_concat_override);
+ const bool batch_cfg_ok = config_->params.batched_cfg &&
+ sd_version_is_unet(version) &&
+ !uncond.empty() &&
+ img_uncond.empty() &&
+ !skip_uncond &&
+ !cache_runtime.ucache_enabled() &&
+ !(is_skiplayer_step && slg_uncond) &&
+ ip_adapter_tokens.empty() &&
+ ip_adapter_uncond_tokens.empty() &&
+ !config_->animatediff_loaded &&
+ (noised_input.dim() < 4 || noised_input.shape()[3] <= 1) &&
+ std::none_of(generation_extensions.begin(),
+ generation_extensions.end(),
+ [](const std::shared_ptr<GenerationExtension>& extension) {
+ return extension->is_enabled();
+ });
+
+ if (batch_cfg_ok) {
+ sd::Tensor<float> batched_out = run_batched_condition(*positive_condition, c_concat_override);
+ if (!batched_out.empty() && batched_out.dim() >= 4 && batched_out.shape()[3] == 2) {
+ auto parts = sd::ops::chunk(batched_out, 2, 3);
+ cond_out = std::move(parts[0]);
+ uncond_out = std::move(parts[1]);
+ }
+ }
+
if (cond_out.empty()) {
- return {};
+ cond_out = run_condition(*positive_condition, c_concat_override);
+ if (cond_out.empty()) {
+ return {};
+ }
}
- if (!uncond.empty()) {
+ if (uncond_out.empty() && !uncond.empty()) {
if (!skip_uncond) {
const std::vector<int>* uncond_skip_layers = nullptr;
if (is_skiplayer_step && slg_uncond) {
diff --git a/src/stable-diffusion.cpp b/src/stable-diffusion.cpp
index b13c055..44a660b 100644
--- a/src/stable-diffusion.cpp
+++ b/src/stable-diffusion.cpp
@@ -335,6 +335,7 @@ void sd_ctx_params_init(sd_ctx_params_t* sd_ctx_params) {
sd_ctx_params->max_vram = nullptr;
sd_ctx_params->disable_prefetch = false;
sd_ctx_params->disable_segmented_compute = false;
+ sd_ctx_params->batched_cfg = true;
sd_ctx_params->eager_load = false;
sd_ctx_params->enable_mmap = false;
sd_ctx_params->diffusion_flash_attn = false;