diff --git a/__init__.py b/__init__.py index 5a60cf7..e736ac4 100644 --- a/__init__.py +++ b/__init__.py @@ -1 +1 @@ -from .conditioning_math import ConditioningMathInvocation, NormalizeConditioningInvocation \ No newline at end of file +from .conditioning_math import ConditioningMathInvocation, NormalizeConditioningInvocation diff --git a/conditioning_math.py b/conditioning_math.py index 1e9dedd..1a7211a 100644 --- a/conditioning_math.py +++ b/conditioning_math.py @@ -1,22 +1,25 @@ -from typing import Literal, Protocol, Optional +from typing import Literal, Optional, Protocol import numpy as np import sympy import torch from invokeai.app.invocations.fields import ( - FluxConditioningField, CogView4ConditioningField, + FluxConditioningField, FluxReduxConditioningField, TensorField, ) from invokeai.app.invocations.flux_redux import FluxReduxOutput from invokeai.app.invocations.primitives import ( - FluxConditioningOutput, CogView4ConditioningOutput, + FluxConditioningOutput, +) +from invokeai.backend.stable_diffusion.diffusion.conditioning_data import ( + CogView4ConditioningInfo, + FLUXConditioningInfo, + SD3ConditioningInfo, ) -from invokeai.backend.stable_diffusion.diffusion.conditioning_data import FLUXConditioningInfo, \ - SD3ConditioningInfo, CogView4ConditioningInfo from invokeai.invocation_api import ( BaseInvocation, BaseInvocationOutput, @@ -33,6 +36,7 @@ invocation, invocation_output, ) + from . import torch_funcs CONDITIONING_OPERATIONS = Literal[ @@ -59,7 +63,10 @@ def apply_operation(operation: CONDITIONING_OPERATIONS, a: torch.Tensor, b: torc if b is None: b = torch.zeros_like(a) original_dtype = a.dtype - a, b, = a.to(dtype=torch.float32), b.to(dtype=torch.float32) + ( + a, + b, + ) = a.to(dtype=torch.float32), b.to(dtype=torch.float32) embeds: torch.Tensor = torch.zeros_like(a) if operation != "APPEND" and a.shape != b.shape: @@ -84,10 +91,12 @@ def apply_operation(operation: CONDITIONING_OPERATIONS, a: torch.Tensor, b: torc class NamedConditioningField(Protocol): - conditioning_name:str + conditioning_name: str -def _load_conditioning(context: InvocationContext, field: NamedConditioningField) -> ( +def _load_conditioning( + context: InvocationContext, field: NamedConditioningField +) -> ( BasicConditioningInfo | SDXLConditioningInfo | FLUXConditioningInfo @@ -110,10 +119,10 @@ def _load_conditioning(context: InvocationContext, field: NamedConditioningField ) class ConditioningMathInvocation(BaseInvocation): """Compute between two conditioning latents""" - + a: ConditioningField = InputField( description="Conditioning A", - input=Input.Connection, #A is required for extra information in some operations + input=Input.Connection, # A is required for extra information in some operations ui_order=0, ) b: Optional[ConditioningField] = InputField( @@ -128,12 +137,13 @@ class ConditioningMathInvocation(BaseInvocation): ui_order=3, ) operation: CONDITIONING_OPERATIONS = InputField( - default="LERP", description="The operation to perform", ui_choice_labels=CONDITIONING_OPERATIONS_LABELS, + default="LERP", + description="The operation to perform", + ui_choice_labels=CONDITIONING_OPERATIONS_LABELS, input=Input.Direct, ui_order=2, ) - @torch.inference_mode() def invoke(self, context: InvocationContext) -> ConditioningOutput: self.check_matching_type(context) @@ -152,16 +162,14 @@ def invoke(self, context: InvocationContext) -> ConditioningOutput: conditioning_info = SDXLConditioningInfo( embeds=embeds, pooled_embeds=pooled_embeds, - add_time_ids=conditioning_A.add_time_ids, #always from A, just includes size information + add_time_ids=conditioning_A.add_time_ids, # always from A, just includes size information ) else: conditioning_info = BasicConditioningInfo(embeds=embeds) conditioning_data = ConditioningFieldData(conditionings=[conditioning_info]) conditioning_name = context.conditioning.save(conditioning_data) - return ConditioningOutput( - conditioning=ConditioningField(conditioning_name=conditioning_name) - ) + return ConditioningOutput(conditioning=ConditioningField(conditioning_name=conditioning_name)) def check_matching_type(self, context): if self.b is None: @@ -185,7 +193,7 @@ def check_matching_type(self, context): class FluxConditioningMathInvocation(ConditioningMathInvocation): a: FluxConditioningField = InputField( description="Conditioning A", - input=Input.Connection, #A is required for extra information in some operations + input=Input.Connection, # A is required for extra information in some operations ui_order=0, ) b: Optional[FluxConditioningField] = InputField( @@ -206,9 +214,7 @@ def invoke(self, context: InvocationContext) -> FluxConditioningOutput: conditioning_info = FLUXConditioningInfo(clip_embeds, t5_embeds) conditioning_data = ConditioningFieldData(conditionings=[conditioning_info]) conditioning_name = context.conditioning.save(conditioning_data) - return FluxConditioningOutput( - conditioning=FluxConditioningField(conditioning_name=conditioning_name) - ) + return FluxConditioningOutput(conditioning=FluxConditioningField(conditioning_name=conditioning_name)) @invocation( @@ -224,24 +230,29 @@ class FluxConditioningFreeformMathInvocation(BaseInvocation): title="c1", input=Input.Connection, ) - c2: FluxConditioningField | None = InputField( - description="Conditioning 2", title="c2", default=None - ) + c2: FluxConditioningField | None = InputField(description="Conditioning 2", title="c2", default=None) c3: FluxConditioningField | None = InputField( - description="Conditioning 3", title="c3", default=None, + description="Conditioning 3", + title="c3", + default=None, ) c4: FluxConditioningField | None = InputField( - description="Conditioning 4", title="c4", default=None, + description="Conditioning 4", + title="c4", + default=None, ) c5: FluxConditioningField | None = InputField( - description="Conditioning 5", title="c5", default=None, + description="Conditioning 5", + title="c5", + default=None, ) a: float = InputField(default=1, title="a") b: float = InputField(default=0, title="b") - formula: str = InputField(description="Formula to apply to conditionings c1–c5. proj, perp, and all torch functions available.", - default="c1") - + formula: str = InputField( + description="Formula to apply to conditionings c1–c5. proj, perp, and all torch functions available.", + default="c1", + ) def invoke(self, context: InvocationContext) -> FluxConditioningOutput: func = self._func_from_string(self.formula) @@ -263,9 +274,7 @@ def invoke(self, context: InvocationContext) -> FluxConditioningOutput: conditioning_info = FLUXConditioningInfo(clip_embeds, t5_embeds) conditioning_data = ConditioningFieldData(conditionings=[conditioning_info]) conditioning_name = context.conditioning.save(conditioning_data) - return FluxConditioningOutput( - conditioning=FluxConditioningField(conditioning_name=conditioning_name) - ) + return FluxConditioningOutput(conditioning=FluxConditioningField(conditioning_name=conditioning_name)) def _clip_embeds( self, context: InvocationContext, field: FluxConditioningField | None, like: torch.Tensor | None = None @@ -308,24 +317,29 @@ class FluxReduxConditioningFreeformMathInvocation(BaseInvocation): title="c1", input=Input.Connection, ) - c2: FluxReduxConditioningField | None = InputField( - description="Conditioning 2", title="c2", default=None - ) + c2: FluxReduxConditioningField | None = InputField(description="Conditioning 2", title="c2", default=None) c3: FluxReduxConditioningField | None = InputField( - description="Conditioning 3", title="c3", default=None, + description="Conditioning 3", + title="c3", + default=None, ) c4: FluxReduxConditioningField | None = InputField( - description="Conditioning 4", title="c4", default=None, + description="Conditioning 4", + title="c4", + default=None, ) c5: FluxReduxConditioningField | None = InputField( - description="Conditioning 5", title="c5", default=None, + description="Conditioning 5", + title="c5", + default=None, ) a: float = InputField(default=1, title="a") b: float = InputField(default=0, title="b") - formula: str = InputField(description="Formula to apply to conditionings c1–c5. proj, perp, and all torch functions available.", - default="c1") - + formula: str = InputField( + description="Formula to apply to conditionings c1–c5. proj, perp, and all torch functions available.", + default="c1", + ) def invoke(self, context: InvocationContext) -> FluxReduxOutput: func = self._func_from_string(self.formula) @@ -342,7 +356,6 @@ def invoke(self, context: InvocationContext) -> FluxReduxOutput: redux_cond=FluxReduxConditioningField(conditioning=TensorField(tensor_name=conditioning_name)) ) - def _redux_embeds( self, context: InvocationContext, field: FluxReduxConditioningField | None, like: torch.Tensor | None = None ) -> torch.Tensor | None: @@ -352,7 +365,6 @@ def _redux_embeds( return torch.zeros_like(like) return None - def _func_from_string(self, formula: str): return sympy.lambdify( sympy.symbols("c1 c2 c3 c4 c5 a b"), @@ -371,7 +383,7 @@ def _func_from_string(self, formula: str): class CogView4ConditioningMathInvocation(ConditioningMathInvocation): a: CogView4ConditioningField = InputField( description="Conditioning A", - input=Input.Connection, #A is required for extra information in some operations + input=Input.Connection, # A is required for extra information in some operations ui_order=0, ) b: Optional[CogView4ConditioningField] = InputField( @@ -389,9 +401,7 @@ def invoke(self, context: InvocationContext) -> CogView4ConditioningOutput: conditioning_info = CogView4ConditioningInfo(glm_embeds) conditioning_data = ConditioningFieldData(conditionings=[conditioning_info]) conditioning_name = context.conditioning.save(conditioning_data) - return CogView4ConditioningOutput( - conditioning=CogView4ConditioningField(conditioning_name=conditioning_name) - ) + return CogView4ConditioningOutput(conditioning=CogView4ConditioningField(conditioning_name=conditioning_name)) @invocation_output("extended_conditioning_output") @@ -405,7 +415,6 @@ class ExtendedConditioningOutput(BaseInvocationOutput): token_space: int = OutputField(description="Number of tokens in the conditioning") - NORMALIZE_OPERATIONS = Literal[ "INFO", "MEAN", @@ -431,24 +440,19 @@ class ExtendedConditioningOutput(BaseInvocationOutput): ) class NormalizeConditioningInvocation(BaseInvocation): """Normalize a conditioning (SD1.5) latent to have a mean and variance similar to another conditioning latent""" - + conditioning: ConditioningField = InputField( description="Conditioning", input=Input.Connection, ) operation: NORMALIZE_OPERATIONS = InputField( - default="INFO", description="The operation to perform", ui_choice_labels=NORMALIZE_OPERATIONS_LABELS, - input=Input.Direct - ) - mean: float = InputField( - default=-0.1, - description="Mean to normalize to" - ) - var: float = InputField( - default=1.0, - description="Standard Deviation to normalize to", - title="Variance" + default="INFO", + description="The operation to perform", + ui_choice_labels=NORMALIZE_OPERATIONS_LABELS, + input=Input.Direct, ) + mean: float = InputField(default=-0.1, description="Mean to normalize to") + var: float = InputField(default=1.0, description="Standard Deviation to normalize to", title="Variance") @torch.no_grad() def invoke(self, context: InvocationContext) -> ExtendedConditioningOutput: @@ -465,12 +469,10 @@ def invoke(self, context: InvocationContext) -> ExtendedConditioningOutput: c = ((c - mean_c) * self.var / std_c) + mean_c elif self.operation == "MEAN_VAR": c = ((c - mean_c) * np.sqrt(self.var) / std_c) + self.mean - + mean_out, std_out, var_out = torch.mean(c), torch.std(c), torch.var(c) - conditioning_data = ConditioningFieldData( - conditionings=[BasicConditioningInfo(embeds=c)] - ) + conditioning_data = ConditioningFieldData(conditionings=[BasicConditioningInfo(embeds=c)]) conditioning_name = context.conditioning.save(conditioning_data) @@ -508,7 +510,7 @@ def _randn_like(tensor: torch.Tensor, generator: torch.Generator) -> torch.Tenso clip = self._clip_embeds(context, self.conditioning) t5 = self._t5_embeds(context, self.conditioning) - generator = torch.Generator(device=clip.device) # Assume this is the same as t5.device + generator = torch.Generator(device=clip.device) # Assume this is the same as t5.device generator.manual_seed(int(self.seed)) @@ -518,9 +520,7 @@ def _randn_like(tensor: torch.Tensor, generator: torch.Generator) -> torch.Tenso conditioning_info = FLUXConditioningInfo(clip_embeds, t5_embeds) conditioning_data = ConditioningFieldData(conditionings=[conditioning_info]) conditioning_name = context.conditioning.save(conditioning_data) - return FluxConditioningOutput( - conditioning=FluxConditioningField(conditioning_name=conditioning_name) - ) + return FluxConditioningOutput(conditioning=FluxConditioningField(conditioning_name=conditioning_name)) def _clip_embeds( self, context: InvocationContext, field: FluxConditioningField | None, like: torch.Tensor | None = None diff --git a/torch_funcs.py b/torch_funcs.py index 1d5d81e..98431ee 100644 --- a/torch_funcs.py +++ b/torch_funcs.py @@ -24,7 +24,7 @@ def proj(a, b): return (torch.mul(a, b).sum() / (torch.norm(b) ** 2)) * b -def slerp(v0: torch.Tensor, v1: torch.Tensor, t: float, *, no_NaN = True, DOT_THRESHOLD: float = 0.9995): +def slerp(v0: torch.Tensor, v1: torch.Tensor, t: float, *, no_NaN=True, DOT_THRESHOLD: float = 0.9995): """Spherical linear interpolation. :param v0: The starting vector @@ -53,9 +53,7 @@ def slerp(v0: torch.Tensor, v1: torch.Tensor, t: float, *, no_NaN = True, DOT_TH can_slerp = ~gotta_lerp t_batch_dim_count: int = max(0, t.dim() - v0.dim()) if isinstance(t, torch.Tensor) else 0 - t_batch_dims: torch.Size = ( - t.shape[:t_batch_dim_count] if isinstance(t, torch.Tensor) else torch.Size([]) - ) + t_batch_dims: torch.Size = t.shape[:t_batch_dim_count] if isinstance(t, torch.Tensor) else torch.Size([]) out = torch.zeros_like(v0.expand(*t_batch_dims, *[-1] * v0.dim())) # if no elements are lerpable, our vectors become 0-dimensional, preventing broadcasting