Bug fixes

More expression tensor operations
Make the return expression handler actually work
This commit is contained in:
blepping
2026-07-11 10:44:15 -06:00
parent ee59df94e3
commit 5ea059bed5
6 changed files with 158 additions and 10 deletions
+1 -1
View File
@@ -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(
+5 -1
View File
@@ -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)
+3 -1
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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)