| import math |
| from dataclasses import dataclass, field |
|
|
| import torch |
|
|
| from .protocols import GuiderProtocol |
|
|
|
|
| @dataclass(frozen=True) |
| class CFGGuider(GuiderProtocol): |
| """ |
| Classifier-free guidance (CFG) guider. |
| Computes the guidance delta as (scale - 1) * (cond - uncond), steering the |
| denoising process toward the conditioned prediction. |
| Attributes: |
| scale: Guidance strength. 1.0 means no guidance, higher values increase |
| adherence to the conditioning. |
| """ |
|
|
| scale: float |
|
|
| def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor: |
| return (self.scale - 1) * (cond - uncond) |
|
|
| def enabled(self) -> bool: |
| return self.scale != 1.0 |
|
|
|
|
| @dataclass(frozen=True) |
| class CFGStarRescalingGuider(GuiderProtocol): |
| """ |
| Calculates the CFG delta between conditioned and unconditioned samples. |
| To minimize offset in the denoising direction and move mostly along the |
| conditioning axis within the distribution, the unconditioned sample is |
| rescaled in accordance with the norm of the conditioned sample. |
| Attributes: |
| scale (float): |
| Global guidance strength. A value of 1.0 corresponds to no extra |
| guidance beyond the base model prediction. Values > 1.0 increase |
| the influence of the conditioned sample relative to the |
| unconditioned one. |
| """ |
|
|
| scale: float |
|
|
| def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor: |
| rescaled_neg = projection_coef(cond, uncond) * uncond |
| return (self.scale - 1) * (cond - rescaled_neg) |
|
|
| def enabled(self) -> bool: |
| return self.scale != 1.0 |
|
|
|
|
| @dataclass(frozen=True) |
| class STGGuider(GuiderProtocol): |
| """ |
| Calculates the STG delta between conditioned and perturbed denoised samples. |
| Perturbed samples are the result of the denoising process with perturbations, |
| e.g. attentions acting as passthrough for certain layers and modalities. |
| Attributes: |
| scale (float): |
| Global strength of the STG guidance. A value of 0.0 disables the |
| guidance. Larger values increase the correction applied in the |
| direction of (pos_denoised - perturbed_denoised). |
| """ |
|
|
| scale: float |
|
|
| def delta(self, pos_denoised: torch.Tensor, perturbed_denoised: torch.Tensor) -> torch.Tensor: |
| return self.scale * (pos_denoised - perturbed_denoised) |
|
|
| def enabled(self) -> bool: |
| return self.scale != 0.0 |
|
|
|
|
| @dataclass(frozen=True) |
| class LtxAPGGuider(GuiderProtocol): |
| """ |
| Calculates the APG (adaptive projected guidance) delta between conditioned |
| and unconditioned samples. |
| To minimize offset in the denoising direction and move mostly along the |
| conditioning axis within the distribution, the (cond - uncond) delta is |
| decomposed into components parallel and orthogonal to the conditioned |
| sample. The `eta` parameter weights the parallel component, while `scale` |
| is applied to the orthogonal component. Optionally, a norm threshold can |
| be used to suppress guidance when the magnitude of the correction is small. |
| Attributes: |
| scale (float): |
| Strength applied to the component of the guidance that is orthogonal |
| to the conditioned sample. Controls how aggressively we move in |
| directions that change semantics but stay consistent with the |
| conditioning manifold. |
| eta (float): |
| Weight of the component of the guidance that is parallel to the |
| conditioned sample. A value of 1.0 keeps the full parallel |
| component; values in [0, 1] attenuate it, and values > 1.0 amplify |
| motion along the conditioning direction. |
| norm_threshold (float): |
| Minimum L2 norm of the guidance delta below which the guidance |
| can be reduced or ignored (depending on implementation). |
| This is useful for avoiding noisy or unstable updates when the |
| guidance signal is very small. |
| """ |
|
|
| scale: float |
| eta: float = 1.0 |
| norm_threshold: float = 0.0 |
|
|
| def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor: |
| guidance = cond - uncond |
| if self.norm_threshold > 0: |
| ones = torch.ones_like(guidance) |
| guidance_norm = guidance.norm(p=2, dim=[-1, -2, -3], keepdim=True) |
| scale_factor = torch.minimum(ones, self.norm_threshold / guidance_norm) |
| guidance = guidance * scale_factor |
| proj_coeff = projection_coef(guidance, cond) |
| g_parallel = proj_coeff * cond |
| g_orth = guidance - g_parallel |
| g_apg = g_parallel * self.eta + g_orth |
|
|
| return g_apg * (self.scale - 1) |
|
|
| def enabled(self) -> bool: |
| return self.scale != 1.0 |
|
|
|
|
| @dataclass(frozen=False) |
| class LegacyStatefulAPGGuider(GuiderProtocol): |
| """ |
| Calculates the APG (adaptive projected guidance) delta between conditioned |
| and unconditioned samples. |
| To minimize offset in the denoising direction and move mostly along the |
| conditioning axis within the distribution, the (cond - uncond) delta is |
| decomposed into components parallel and orthogonal to the conditioned |
| sample. The `eta` parameter weights the parallel component, while `scale` |
| is applied to the orthogonal component. Optionally, a norm threshold can |
| be used to suppress guidance when the magnitude of the correction is small. |
| Attributes: |
| scale (float): |
| Strength applied to the component of the guidance that is orthogonal |
| to the conditioned sample. Controls how aggressively we move in |
| directions that change semantics but stay consistent with the |
| conditioning manifold. |
| eta (float): |
| Weight of the component of the guidance that is parallel to the |
| conditioned sample. A value of 1.0 keeps the full parallel |
| component; values in [0, 1] attenuate it, and values > 1.0 amplify |
| motion along the conditioning direction. |
| norm_threshold (float): |
| Minimum L2 norm of the guidance delta below which the guidance |
| can be reduced or ignored (depending on implementation). |
| This is useful for avoiding noisy or unstable updates when the |
| guidance signal is very small. |
| momentum (float): |
| Exponential moving-average coefficient for accumulating guidance |
| over time. running_avg = momentum * running_avg + guidance |
| """ |
|
|
| scale: float |
| eta: float |
| norm_threshold: float = 5.0 |
| momentum: float = 0.0 |
| |
| |
| running_avg: torch.Tensor | None = None |
|
|
| def delta(self, cond: torch.Tensor, uncond: torch.Tensor) -> torch.Tensor: |
| guidance = cond - uncond |
| if self.momentum != 0: |
| if self.running_avg is None: |
| self.running_avg = guidance.clone() |
| else: |
| self.running_avg = self.momentum * self.running_avg + guidance |
| guidance = self.running_avg |
|
|
| if self.norm_threshold > 0: |
| ones = torch.ones_like(guidance) |
| guidance_norm = guidance.norm(p=2, dim=[-1, -2, -3], keepdim=True) |
| scale_factor = torch.minimum(ones, self.norm_threshold / guidance_norm) |
| guidance = guidance * scale_factor |
|
|
| proj_coeff = projection_coef(guidance, cond) |
| g_parallel = proj_coeff * cond |
| g_orth = guidance - g_parallel |
| g_apg = g_parallel * self.eta + g_orth |
|
|
| return g_apg * self.scale |
|
|
| def enabled(self) -> bool: |
| return self.scale != 0.0 |
|
|
|
|
| @dataclass(frozen=True) |
| class MultiModalGuiderParams: |
| cfg_scale: float = 1.0 |
| stg_scale: float = 0.0 |
| stg_blocks: list[int] | None = field(default_factory=list) |
| rescale_scale: float = 0.0 |
| modality_scale: float = 1.0 |
| skip_step: int = 0 |
|
|
|
|
| @dataclass(frozen=True) |
| class MultiModalGuider: |
| params: MultiModalGuiderParams |
| negative_context: torch.Tensor | None = None |
|
|
| def calculate( |
| self, |
| cond: torch.Tensor, |
| uncond_text: torch.Tensor | float, |
| uncond_perturbed: torch.Tensor | float, |
| uncond_modality: torch.Tensor | float, |
| ) -> torch.Tensor: |
| if cond is None: |
| return None |
| if uncond_text is None: |
| uncond_text = cond |
| if uncond_perturbed is None: |
| uncond_perturbed = cond |
| if uncond_modality is None: |
| uncond_modality = cond |
| pred = ( |
| cond |
| + (self.params.cfg_scale - 1) * (cond - uncond_text) |
| + self.params.stg_scale * (cond - uncond_perturbed) |
| + (self.params.modality_scale - 1) * (cond - uncond_modality) |
| ) |
| if self.params.rescale_scale != 0: |
| pred_std = pred.std().clamp_min(1e-6) |
| factor = cond.std() / pred_std |
| factor = self.params.rescale_scale * factor + (1 - self.params.rescale_scale) |
| pred = pred * factor |
| return pred |
|
|
| def do_unconditional_generation(self) -> bool: |
| return not math.isclose(self.params.cfg_scale, 1.0) |
|
|
| def do_perturbed_generation(self) -> bool: |
| return not math.isclose(self.params.stg_scale, 0.0) |
|
|
| def do_isolated_modality_generation(self) -> bool: |
| return not math.isclose(self.params.modality_scale, 1.0) |
|
|
| def should_skip_step(self, step: int) -> bool: |
| if self.params.skip_step == 0: |
| return False |
| return step % (self.params.skip_step + 1) != 0 |
|
|
|
|
| def projection_coef(to_project: torch.Tensor, project_onto: torch.Tensor) -> torch.Tensor: |
| batch_size = to_project.shape[0] |
| positive_flat = to_project.reshape(batch_size, -1) |
| negative_flat = project_onto.reshape(batch_size, -1) |
| dot_product = torch.sum(positive_flat * negative_flat, dim=1, keepdim=True) |
| squared_norm = torch.sum(negative_flat**2, dim=1, keepdim=True) + 1e-8 |
| return dot_product / squared_norm |
|
|