diff --git a/py/custom_noise/nodes.py b/py/custom_noise/nodes.py index 670fb07..631a799 100644 --- a/py/custom_noise/nodes.py +++ b/py/custom_noise/nodes.py @@ -1635,7 +1635,7 @@ class SamplerNodeConfigOverride(metaclass=IntegratedNode): **kwargs: dict, ) -> torch.Tensor: nonlocal ref_latent, filter_refs, immiscible_counter, noise_prev - if not isigma_end <= s.max() <= isigma_start: + if not isigma_end <= s.max() <= isigma_start or icfg.size == 0: return override_noise_sampler(s, sn, *args, **kwargs) if icfg.filter_noise or icfg.filter_result: curr_refs = filtering.FilterRefs( diff --git a/py/expression/expression.py b/py/expression/expression.py index 474563a..27df9ff 100644 --- a/py/expression/expression.py +++ b/py/expression/expression.py @@ -13,6 +13,7 @@ from .types import ( ExpKV, ExpMethodAp, ExpOp, + ExpReturn, ExpStatements, ExpSym, ExpTuple, @@ -65,7 +66,10 @@ class Expression: tqdm.write(f"* OCS: EVAL: {self.expr}") if not isinstance(self.expr, ExpBase): return self.expr - return self.expr.eval(handlers, *args, **kwargs) + try: + return self.expr.eval(handlers, *args, **kwargs) + except ExpReturn as expret: + return expret.args[0] def __len__(self): return len(self.expr) diff --git a/py/expression/handler.py b/py/expression/handler.py index e0287f8..851c650 100644 --- a/py/expression/handler.py +++ b/py/expression/handler.py @@ -62,10 +62,12 @@ class BaseHandler: def __call__(self, obj, *, getter): try: val = self.handle(obj, getter) - return self.validate_output(obj, val) + except ExpReturn: + raise except Exception as exc: tb = traceback.format_exc() raise HandlerError(f'Error evaluating "{obj.name}": {exc!s}\n{tb}') from exc + return self.validate_output(obj, val) def safe_get(self, key, obj, getter=None, *, default=Empty): str_key = isinstance(key, str) diff --git a/py/expression_handlers.py b/py/expression_handlers.py index 6cec4fc..47def1d 100644 --- a/py/expression_handlers.py +++ b/py/expression_handlers.py @@ -8,7 +8,7 @@ import torch from . import expression as expr from . import latent, unsafe_expression_whitelists from .external import MODULES as EXT -from .latent import OCSTAESD, ImageBatch, normalize_to_scale +from .latent import OCSTAESD, ImageBatch, flip_tensor_range, normalize_to_scale from .utils import quantile_normalize, resolve_value, scale_noise ALLOW_UNSAFE = os.environ.get("COMFYUI_OCS_ALLOW_UNSAFE_EXPRESSIONS") is not None @@ -307,6 +307,102 @@ class NewFullHandler(NormHandler): return tensor.new_full(shape, value) +class InvertRangeHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("dim"), + ) + + def handle(self, obj, getter): + tensor, dim = self.safe_get_all(obj, getter) + if dim < 0: + dim += tensor.ndim + if dim < 0 or dim >= tensor.ndim: + raise ValueError( + f"Dimension out of range, wanted {dim}, tensor has {tensor.ndim} dimension(s)" + ) + return flip_tensor_range(tensor, dim=dim) + + +class MinHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("dim"), + ) + + def handle(self, obj, getter): + tensor, dim = self.safe_get_all(obj, getter) + return tensor.min(dim=dim, keepdim=True).values + + +class MaxHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("dim"), + ) + + def handle(self, obj, getter): + tensor, dim = self.safe_get_all(obj, getter) + return tensor.max(dim=dim, keepdim=True).values + + +class CumSumHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("dim"), + ) + + def handle(self, obj, getter): + tensor, dim = self.safe_get_all(obj, getter) + return tensor.cumsum(dim=dim) + + +class MinimumHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor1"), + expr.Arg.tensor("tensor2"), + ) + + def handle(self, obj, getter): + tensor1, tensor2 = self.safe_get_all(obj, getter) + return tensor1.minimum(tensor2) + + +class MaximumHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor1"), + expr.Arg.tensor("tensor2"), + ) + + def handle(self, obj, getter): + tensor1, tensor2 = self.safe_get_all(obj, getter) + return tensor1.maximum(tensor2) + + +class MoveDimHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("from_dim"), + expr.Arg.integer("to_dim", -1), + ) + + def handle(self, obj, getter): + tensor, from_dim, to_dim = self.safe_get_all(obj, getter) + return tensor.movedim(from_dim, to_dim) + + +class FlattenHandler(NormHandler): + input_validators = ( + expr.Arg.tensor("tensor"), + expr.Arg.integer("start_dim", 0), + expr.Arg.integer("end_dim", -1), + ) + + def handle(self, obj, getter): + tensor, start_dim, end_dim = self.safe_get_all(obj, getter) + return tensor.flatten(start_dim=start_dim, end_dim=end_dim) + + class BlendHandler(NormHandler): input_validators = ( expr.Arg.tensor("tensor1"), @@ -699,10 +795,18 @@ TENSOR_OP_HANDLERS = { "t_clone": CloneHandler(), "t_newfull": NewFullHandler(), "t_copysign": CopySignHandler(), + "t_flatten": FlattenHandler(), + "t_movedim": MoveDimHandler(), + "t_min": MinHandler(), + "t_max": MaxHandler(), + "t_minimum": MinimumHandler(), + "t_maximum": MaximumHandler(), + "t_cumsum": CumSumHandler(), "t_contrast_adaptive_sharpening": ContrastAdaptiveSharpeningHandler(), "t_scale": ScaleHandler(), "t_noise": NoiseHandler(), "t_shape": ShapeHandler(), + "t_invert_range": InvertRangeHandler(), "t_gaussianblur2d": GaussianBlur2DHandler(), "t_rgb_latent": RGBLatentHandler(), "t_snf_guidance": SNFGuidanceHandler(), diff --git a/py/latent.py b/py/latent.py index 22738cc..55a2c18 100644 --- a/py/latent.py +++ b/py/latent.py @@ -1,13 +1,11 @@ +import folder_paths +import latent_preview import numpy as np import torch import torch.nn.functional as F - -import folder_paths -import latent_preview - +from comfy import latent_formats from comfy.taesd.taesd import TAESD from comfy.utils import bislerp -from comfy import latent_formats from .external import MODULES as EXT @@ -125,6 +123,45 @@ def contrast_adaptive_sharpening( # noqa: PLR0914 return output.reshape(*orig_shape) +def flip_tensor_range( + x: torch.Tensor, + *, + min_neg: torch.Tensor | None = None, + max_pos: torch.Tensor | None = None, + return_ranges: bool = False, + dim: int = -1, + eps: float | None = None, +) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + if eps is None: + eps = torch.finfo(x.dtype).eps * 1.25 + # 1. Use the provided maximum positive values, or calculate them dynamically + if max_pos is None: + max_pos = ( + torch.clamp_min(x, 0.0).max(dim=dim, keepdim=True).values.clamp_min_(eps) + ) + + # 2. Use the provided minimum negative values, or calculate them dynamically + if min_neg is None: + min_neg = ( + torch.clamp_max(x, 0.0).min(dim=dim, keepdim=True).values.clamp_max_(-eps) + ) + + # 3. Separate positive and negative elements + is_pos = x >= 0 + + # 4. Flip positive side: [0, max_pos] -> [eps, max_pos + eps] + x_pos = x.clamp_min(eps) + flipped_pos = (max_pos + eps) - x_pos + + # 5. Flip negative side: [min_neg, 0] -> [min_neg - eps, -eps] + x_neg = x.clamp_max(-eps) + flipped_neg = (min_neg - eps) - x_neg + + # 6. Recombine the domains + result = torch.where(is_pos, flipped_pos, flipped_neg) + return (result, max_pos, min_neg) if return_ranges else result + + class ImageBatch(tuple): __slots__ = () diff --git a/py/nodes.py b/py/nodes.py index c1ab792..2653b7c 100644 --- a/py/nodes.py +++ b/py/nodes.py @@ -800,7 +800,8 @@ class ExpressionFilteredLatentOperation: **kwargs: dict, ) -> torch.Tensor: refs = FilterRefs( - kvs={ + kvs=kwargs + | { "sigma": sigma.clone() if isinstance(sigma, torch.Tensor) else sigma, "sigma_float": sigma.max().item() if isinstance(sigma, torch.Tensor)