Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion __init__.py
Original file line number Diff line number Diff line change
@@ -1 +1 @@
from .conditioning_math import ConditioningMathInvocation, NormalizeConditioningInvocation
from .conditioning_math import ConditioningMathInvocation, NormalizeConditioningInvocation
136 changes: 68 additions & 68 deletions conditioning_math.py
Original file line number Diff line number Diff line change
@@ -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,
Expand All @@ -33,6 +36,7 @@
invocation,
invocation_output,
)

from . import torch_funcs

CONDITIONING_OPERATIONS = Literal[
Expand All @@ -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:
Expand All @@ -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
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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:
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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
Expand Down Expand Up @@ -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)
Expand All @@ -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:
Expand All @@ -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"),
Expand All @@ -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(
Expand All @@ -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")
Expand All @@ -405,7 +415,6 @@ class ExtendedConditioningOutput(BaseInvocationOutput):
token_space: int = OutputField(description="Number of tokens in the conditioning")



NORMALIZE_OPERATIONS = Literal[
"INFO",
"MEAN",
Expand All @@ -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:
Expand All @@ -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)

Expand Down Expand Up @@ -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))

Expand All @@ -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
Expand Down
6 changes: 2 additions & 4 deletions torch_funcs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand Down