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;