Bug fixes
More expression tensor operations Make the return expression handler actually work
This commit is contained in:
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
+105
-1
@@ -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(),
|
||||
|
||||
+42
-5
@@ -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__ = ()
|
||||
|
||||
|
||||
+2
-1
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user