Repository navigation
feat: run conditional and unconditional CFG in one batched UNet forward #2085
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
f4e625c
811b4b0
879be0a
11bff8d
b436b68
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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]); | ||
| } | ||
| } | ||
|
Comment on lines
+2736
to
+2742
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🚀 Performance & Scalability | 🟡 Minor | ⚡ Quick win A silent batched-path failure doubles the per-step cost.
After the first failure, latch a per- Proposed fix- if (batch_cfg_ok) {
+ if (batch_cfg_ok && !batched_cfg_failed) {
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) {
...
+ } else {
+ batched_cfg_failed = true;
+ LOG_WARN("batched CFG unavailable; falling back to separate forwards");
}
}Declare 🤖 Prompt for AI Agents |
||
|
|
||
| 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) { | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🚀 Performance & Scalability | 🟡 Minor | ⚡ Quick win
🔎 Supported by static analysis
🏁 Script executed:
Repository: leejet/stable-diffusion.cpp
Length of output: 9486
Reuse unconditional controls in the fallback path.
When the batched path fails or returns an output with an invalid shape,
run_condition(uncond, ...)computes the unconditional ControlNet controls again becauseuncond_controls_readyremainsfalse. This duplicate pass occurs only on fallback steps, not on every step. Preserve the controls computed byrun_batched_conditionand mark them ready before the unconditional fallback.🤖 Prompt for AI Agents