From f900a6457ae2713ab3bacc9b2b2ba957e5313a38 Mon Sep 17 00:00:00 2001 From: "remyx-ai[bot]" <289541483+remyx-ai[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:17:14 +0000 Subject: [PATCH 1/2] Add semantic-aware CFG guider with per-region scale map --- src/diffusers/guiders/__init__.py | 1 + .../guiders/semantic_aware_guidance.py | 209 ++++++++++++++++++ tests/guiders/__init__.py | 0 tests/guiders/test_semantic_aware_guidance.py | 77 +++++++ 4 files changed, 287 insertions(+) create mode 100644 src/diffusers/guiders/semantic_aware_guidance.py create mode 100644 tests/guiders/__init__.py create mode 100644 tests/guiders/test_semantic_aware_guidance.py diff --git a/src/diffusers/guiders/__init__.py b/src/diffusers/guiders/__init__.py index b6653817dc95..1c78782474a3 100644 --- a/src/diffusers/guiders/__init__.py +++ b/src/diffusers/guiders/__init__.py @@ -26,6 +26,7 @@ from .guider_utils import BaseGuidance from .magnitude_aware_guidance import MagnitudeAwareGuidance from .perturbed_attention_guidance import PerturbedAttentionGuidance + from .semantic_aware_guidance import SemanticAwareGuidance from .skip_layer_guidance import SkipLayerGuidance from .smoothed_energy_guidance import SmoothedEnergyGuidance from .tangential_classifier_free_guidance import TangentialClassifierFreeGuidance diff --git a/src/diffusers/guiders/semantic_aware_guidance.py b/src/diffusers/guiders/semantic_aware_guidance.py new file mode 100644 index 000000000000..b2672cf94ca8 --- /dev/null +++ b/src/diffusers/guiders/semantic_aware_guidance.py @@ -0,0 +1,209 @@ +# Copyright 2025 The HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +from typing import TYPE_CHECKING + +import torch +import torch.nn.functional as F + +from ..configuration_utils import register_to_config +from .guider_utils import BaseGuidance, GuiderOutput, rescale_noise_cfg + + +if TYPE_CHECKING: + from ..modular_pipelines.modular_pipeline import BlockState + + +class SemanticAwareGuidance(BaseGuidance): + """ + Semantic-aware Classifier-Free Guidance (S-CFG): https://huggingface.co/papers/2404.05384 + + A global CFG scale applies the same guidance strength everywhere, so regions whose text guidance already produces a + large update get over-guided while weak regions stay under-guided -- the "spatial inconsistency" the paper targets. + S-CFG replaces the scalar with a per-region scale map that equalizes the guidance strength across the latent, then + applies `pred = pred_uncond + scale_map * (pred_cond - pred_uncond)` (the standard guider contract). + + **Adaptation note.** The paper derives its semantic regions from cross-/self-attention maps collected via per-model + attention hooks. That segmentation is an auxiliary signal that is not available at the guider contract (which only + receives `pred_cond` / `pred_uncond`) and is architecture-specific. This implementation keeps the paper's core + mechanism -- a spatially varying CFG scale map that equalizes per-region guidance strength -- at full fidelity, but + substitutes the attention-based segmentation with a parameter-free proxy: the local guidance magnitude computed + directly from `pred_cond - pred_uncond` and pooled over a `window_size` neighborhood to reach region-level (rather + than pixel-level) granularity. Each position is rescaled toward the sample's mean guidance strength, so weak regions + are boosted and strong regions are damped, exactly as in S-CFG. When the predictions are not spatial (`ndim != 4`, + e.g. sequence-shaped transformer outputs) the guider falls back to standard scalar CFG. + + Args: + guidance_scale (`float`, defaults to `7.5`): + The reference CFG scale. The per-region scale map is centered on this value and rescaled around it to + equalize semantic strengths. Higher values give stronger prompt conditioning. + window_size (`int`, defaults to `3`): + Odd side length of the square neighborhood used to pool the local guidance magnitude into a region-level + estimate. `1` disables pooling (pixel-level); larger values approximate coarser semantic regions. + rescale_clamp (`float`, defaults to `2.0`): + Maximum multiplicative deviation of the per-region scale from `guidance_scale`. The scale map is clamped to + `[guidance_scale / rescale_clamp, guidance_scale * rescale_clamp]` to keep low-magnitude regions from + producing runaway guidance. + guidance_rescale (`float`, defaults to `0.0`): + The rescale factor applied to the noise predictions. This is used to improve image quality and fix + overexposure. Based on Section 3.4 from [Common Diffusion Noise Schedules and Sample Steps are + Flawed](https://huggingface.co/papers/2305.08891). + use_original_formulation (`bool`, defaults to `False`): + Whether to use the original formulation of classifier-free guidance as proposed in the paper. By default, + we use the diffusers-native implementation that has been in the codebase for a long time. See + [~guiders.classifier_free_guidance.ClassifierFreeGuidance] for more details. + start (`float`, defaults to `0.0`): + The fraction of the total number of denoising steps after which guidance starts. + stop (`float`, defaults to `1.0`): + The fraction of the total number of denoising steps after which guidance stops. + """ + + _input_predictions = ["pred_cond", "pred_uncond"] + + @register_to_config + def __init__( + self, + guidance_scale: float = 7.5, + window_size: int = 3, + rescale_clamp: float = 2.0, + guidance_rescale: float = 0.0, + use_original_formulation: bool = False, + start: float = 0.0, + stop: float = 1.0, + enabled: bool = True, + ): + super().__init__(start, stop, enabled) + + if window_size < 1 or window_size % 2 == 0: + raise ValueError(f"Expected `window_size` to be a positive odd integer, but got {window_size}.") + if rescale_clamp < 1.0: + raise ValueError(f"Expected `rescale_clamp` to be >= 1.0, but got {rescale_clamp}.") + + self.guidance_scale = guidance_scale + self.window_size = window_size + self.rescale_clamp = rescale_clamp + self.guidance_rescale = guidance_rescale + self.use_original_formulation = use_original_formulation + + def prepare_inputs(self, data: dict[str, tuple[torch.Tensor, torch.Tensor]]) -> list["BlockState"]: + tuple_indices = [0] if self.num_conditions == 1 else [0, 1] + data_batches = [] + for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): + data_batch = self._prepare_batch(data, tuple_idx, input_prediction) + data_batches.append(data_batch) + return data_batches + + def prepare_inputs_from_block_state( + self, data: "BlockState", input_fields: dict[str, str | tuple[str, str]] + ) -> list["BlockState"]: + tuple_indices = [0] if self.num_conditions == 1 else [0, 1] + data_batches = [] + for tuple_idx, input_prediction in zip(tuple_indices, self._input_predictions): + data_batch = self._prepare_batch_from_block_state(input_fields, data, tuple_idx, input_prediction) + data_batches.append(data_batch) + return data_batches + + def forward(self, pred_cond: torch.Tensor, pred_uncond: torch.Tensor | None = None) -> GuiderOutput: + pred = None + + if not self._is_cfg_enabled(): + pred = pred_cond + else: + pred = semantic_aware_guidance( + pred_cond, + pred_uncond, + self.guidance_scale, + self.window_size, + self.rescale_clamp, + self.use_original_formulation, + ) + + if self.guidance_rescale > 0.0: + pred = rescale_noise_cfg(pred, pred_cond, self.guidance_rescale) + + return GuiderOutput(pred=pred, pred_cond=pred_cond, pred_uncond=pred_uncond) + + @property + def is_conditional(self) -> bool: + return self._count_prepared == 1 + + @property + def num_conditions(self) -> int: + num_conditions = 1 + if self._is_cfg_enabled(): + num_conditions += 1 + return num_conditions + + def _is_cfg_enabled(self) -> bool: + if not self._enabled: + return False + + is_within_range = True + if self._num_inference_steps is not None: + skip_start_step = int(self._start * self._num_inference_steps) + skip_stop_step = int(self._stop * self._num_inference_steps) + is_within_range = skip_start_step <= self._step < skip_stop_step + + is_close = False + if self.use_original_formulation: + is_close = math.isclose(self.guidance_scale, 0.0) + else: + is_close = math.isclose(self.guidance_scale, 1.0) + + return is_within_range and not is_close + + +def semantic_scale_map( + diff: torch.Tensor, + guidance_scale: float, + window_size: int = 3, + rescale_clamp: float = 2.0, + eps: float = 1e-4, +) -> torch.Tensor: + """ + Builds the per-region CFG scale map used by S-CFG from the guidance term `diff = pred_cond - pred_uncond`. + + The local guidance strength is the channel-norm of `diff` at each spatial position, pooled over a `window_size` + neighborhood (a parameter-free proxy for the paper's attention-derived semantic regions). Each position is rescaled + toward the per-sample mean strength so weak regions are boosted and strong regions are damped, then clamped to keep + the guidance well-conditioned. Expects a 4D `(B, C, H, W)` tensor. + """ + magnitude = diff.pow(2).sum(dim=1, keepdim=True).clamp_min(0.0).sqrt() + if window_size > 1: + pad = window_size // 2 + magnitude = F.avg_pool2d(magnitude, kernel_size=window_size, stride=1, padding=pad, count_include_pad=False) + reference = magnitude.mean(dim=(2, 3), keepdim=True) + scale_map = guidance_scale * reference / magnitude.clamp_min(eps) + scale_map = scale_map.clamp(guidance_scale / rescale_clamp, guidance_scale * rescale_clamp) + return scale_map + + +def semantic_aware_guidance( + pred_cond: torch.Tensor, + pred_uncond: torch.Tensor, + guidance_scale: float, + window_size: int = 3, + rescale_clamp: float = 2.0, + use_original_formulation: bool = False, +) -> torch.Tensor: + diff = pred_cond - pred_uncond + pred = pred_cond if use_original_formulation else pred_uncond + + if diff.ndim != 4: + # Non-spatial predictions (e.g. sequence-shaped transformer outputs): fall back to scalar CFG. + return pred + guidance_scale * diff + + scale_map = semantic_scale_map(diff, guidance_scale, window_size, rescale_clamp) + return pred + scale_map * diff diff --git a/tests/guiders/__init__.py b/tests/guiders/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/guiders/test_semantic_aware_guidance.py b/tests/guiders/test_semantic_aware_guidance.py new file mode 100644 index 000000000000..81bf545da269 --- /dev/null +++ b/tests/guiders/test_semantic_aware_guidance.py @@ -0,0 +1,77 @@ +# Copyright 2025 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import torch + +# Import the new guider through the existing `diffusers.guiders` package (the call-site wiring), and pull the existing +# ClassifierFreeGuidance from the same package to assert the S-CFG guider is a faithful drop-in extension of it. +from diffusers.guiders import ClassifierFreeGuidance, SemanticAwareGuidance + + +def _cfg_pred(guider, pred_cond, pred_uncond): + return guider.forward(pred_cond, pred_uncond).pred + + +def test_registered_alongside_classifier_free_guidance(): + # Both guiders resolve from the existing package and share the CFG contract. + guider = SemanticAwareGuidance(guidance_scale=7.5) + assert guider._input_predictions == ["pred_cond", "pred_uncond"] + assert guider.num_conditions == 2 + + +def test_reduces_to_scalar_cfg_when_guidance_is_spatially_uniform(): + # When the guidance term has a constant magnitude everywhere, S-CFG's scale map collapses to the scalar + # guidance_scale, so its output must match the existing ClassifierFreeGuidance exactly. + torch.manual_seed(0) + pred_uncond = torch.randn(2, 4, 8, 8) + diff = torch.full((2, 4, 8, 8), 0.5) + pred_cond = pred_uncond + diff + + cfg = ClassifierFreeGuidance(guidance_scale=7.5) + scfg = SemanticAwareGuidance(guidance_scale=7.5) + + torch.testing.assert_close(_cfg_pred(scfg, pred_cond, pred_uncond), _cfg_pred(cfg, pred_cond, pred_uncond)) + + +def test_equalizes_semantic_strength_across_regions(): + # A strong region (large guidance magnitude) should receive a smaller effective scale than a weak region, so that + # the per-region guidance strengths are pulled toward each other -- the core S-CFG behavior. + pred_uncond = torch.zeros(1, 4, 8, 8) + diff = torch.zeros(1, 4, 8, 8) + diff[..., :4] = 2.0 # strong left half + diff[..., 4:] = 0.5 # weak right half + pred_cond = pred_uncond + diff + + guidance_scale = 6.0 + scfg = SemanticAwareGuidance(guidance_scale=guidance_scale, window_size=1, rescale_clamp=4.0) + pred = _cfg_pred(scfg, pred_cond, pred_uncond) + + # Recover the applied scale per region: pred = pred_uncond + scale * diff -> scale = (pred - pred_uncond) / diff. + strong_scale = (pred[..., :4] / diff[..., :4]).mean().item() + weak_scale = (pred[..., 4:] / diff[..., 4:]).mean().item() + + assert strong_scale < guidance_scale < weak_scale + # Effective guidance strengths (scale * magnitude) are closer than the raw magnitudes were. + assert strong_scale * 2.0 - weak_scale * 0.5 < 2.0 - 0.5 + + +def test_falls_back_to_scalar_cfg_for_non_spatial_predictions(): + # Sequence-shaped (3D) predictions have no spatial grid, so S-CFG must behave exactly like scalar CFG. + pred_uncond = torch.randn(2, 16, 32) + pred_cond = pred_uncond + torch.randn(2, 16, 32) + + cfg = ClassifierFreeGuidance(guidance_scale=5.0) + scfg = SemanticAwareGuidance(guidance_scale=5.0) + + torch.testing.assert_close(_cfg_pred(scfg, pred_cond, pred_uncond), _cfg_pred(cfg, pred_cond, pred_uncond)) From 5137bfa37e5ad4ecda029a4d39c26c8f4d38b1cf Mon Sep 17 00:00:00 2001 From: "remyx-ai[bot]" <289541483+remyx-ai[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 14:22:58 +0000 Subject: [PATCH 2/2] chore: align PR with target-repo conventions Convention-shape patches extracted from huggingface/diffusers's recent merged PRs. Algorithm logic is left untouched. Ruff auto-fixed lint-trivial issues on patched files. --- docs/source/en/api/modular_diffusers/guiders.md | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/docs/source/en/api/modular_diffusers/guiders.md b/docs/source/en/api/modular_diffusers/guiders.md index a24eb7220749..53e761a8333c 100644 --- a/docs/source/en/api/modular_diffusers/guiders.md +++ b/docs/source/en/api/modular_diffusers/guiders.md @@ -37,3 +37,7 @@ Guiders are components in Modular Diffusers that control how the diffusion proce ## TangentialClassifierFreeGuidance [[autodoc]] diffusers.guiders.tangential_classifier_free_guidance.TangentialClassifierFreeGuidance + +## SemanticAwareGuidance + +[[autodoc]] diffusers.guiders.semantic_aware_guidance.SemanticAwareGuidance