From 36fcbd48c8db46e49f945dc365c9135b4852960a Mon Sep 17 00:00:00 2001 From: mcDandy Date: Mon, 20 Apr 2026 14:24:37 +0200 Subject: [PATCH] add flow magnitude and angle, histogram --- README.md | 8 +++- more_math/Parser/MathExpr.g4 | 15 +++++-- more_math/Parser/UnifiedMathVisitor.py | 61 +++++++++++++++++++++++++- web/script_text_input.js | 2 +- 4 files changed, 79 insertions(+), 7 deletions(-) diff --git a/README.md b/README.md index ff805ed..e4039bc 100644 --- a/README.md +++ b/README.md @@ -163,6 +163,8 @@ You can also get the node from comfy manager under the name of More math. - `pad(tensor, padding)`: Pads a tensor with specified padding (pair for each dimension). For example, `[1,2,0,0]` adds 1 element before and 2 elements after in the first dimension, and no padding in the second dimension. - `concatenate(tensor1, tensor2, dim)` or `concat` or `cat`: Concatenates two tensors along specified dimension. It things of everything as tensor. - `roll(tensor, shifts, dims)`: Rolls tensor along specified dimensions by given shifts. Elements that roll beyond the last position are re-introduced at the first position. +- `where(cond, a, b)`: Element-wise selection. Returns values from `a` where `cond` is non-zero/true, otherwise from `b`. Supports scalars, lists, and tensors (with tensor broadcasting). +- `histogram(x, bins, min, max)` or `hist`: Computes histogram counts of tensor values in range `[min, max]` using `bins` bins. Returns a 1D tensor of counts. ### Advanced Tensor Operations @@ -204,13 +206,15 @@ You can also get the node from comfy manager under the name of More math. - `motion_mask(flow)`: Generates an occlusion/motion mask from optical flow vectors. - `flow`: Flow vectors [B, H, W, 2]. - **Returns**: Mask [B, H, W] in range [0, 1]. -- `flow_to_image(flow)` or `flow_view(flow)`: Converts flow vectors to an RGB image for visualization. +- `flow_to_image(flow)`: Converts flow vectors to an RGB image for visualization. - `flow`: Flow vectors [B, H, W, 2]. - **Returns**: RGB image [B, H, W, 3]. -- `flow_apply(image, flow)` or `apply_flow(image, flow)`: Warps an image using optical flow vectors. +- `flow_apply(image, flow)`: Warps an image using optical flow vectors. - `image`: Image [B, H, W, C]. - `flow`: Flow vectors [B, H, W, 2] from `rife()`. - **Returns**: Warped image [B, H, W, C]. +- `flow_mag(flow)` or `flow_magnitude(flow)`: Returns optical flow vector magnitude `sqrt(dx^2 + dy^2)`. +- `flow_ang(flow)` or `flow_angle(flow)`: Returns optical flow direction in **radians** using `atan2(dy, dx)`. ### FFT (Tensor Only) diff --git a/more_math/Parser/MathExpr.g4 b/more_math/Parser/MathExpr.g4 index c687a25..5a1cc7e 100644 --- a/more_math/Parser/MathExpr.g4 +++ b/more_math/Parser/MathExpr.g4 @@ -171,7 +171,10 @@ func1: | MORPH_OPEN LPAREN expr (COMMA expr)? RPAREN # MorphOpenFunc | MORPH_CLOSE LPAREN expr (COMMA expr)? RPAREN # MorphCloseFunc | INT LPAREN expr RPAREN # IntFunc - | FLOAT LPAREN expr RPAREN # FloatFunc; + | FLOAT LPAREN expr RPAREN # FloatFunc + | FLOW_MAG LPAREN expr RPAREN # FlowMagFunc + | FLOW_ANG LPAREN expr RPAREN # FlowAngFunc +; func2: POWE LPAREN expr COMMA expr RPAREN # PowFunc @@ -232,13 +235,15 @@ func3: | LINSPACE LPAREN expr COMMA expr COMMA expr RPAREN # LinspaceFunc | ROLL LPAREN expr COMMA expr (COMMA expr)? RPAREN # RollFunc | RGB_TO_HSV LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # RgbToHsvFunc - | HSV_TO_RGB LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # HsvToRgbFunc; + | HSV_TO_RGB LPAREN expr (COMMA expr COMMA expr)? (COMMA expr)? RPAREN # HsvToRgbFunc + | WHERE LPAREN expr COMMA expr COMMA expr RPAREN # WhereFunc; func4: SWAP LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # SwapFunc | NVL LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # NvlFunc | LOGSPACE LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # ĹogspaceFunc - | DIST LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # DistFunc; + | DIST LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # DistFunc + | HISTOGRAM LPAREN expr COMMA expr COMMA expr COMMA expr RPAREN # HistogramFunc; func5: REMAP LPAREN expr COMMA expr COMMA expr COMMA expr COMMA expr RPAREN # RemapFunc; @@ -368,6 +373,10 @@ FLOW_APPLY: 'flow_apply'; BATCH_SHUFFLE: 'batch_shuffle' | 'shuffle' | 'select'; MOTION_MASK: 'motion_mask'; FLOW_TO_IMAGE: 'flow_to_image'; +FLOW_MAG: 'flow_mag' | 'flow_magnitude'; +FLOW_ANG: 'flow_ang' | 'flow_angle'; +WHERE: 'where'; +HISTOGRAM: 'histogram' | 'hist'; OVERLAY: 'overlay'; PAD: 'pad'; CROSS: 'cross'; diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 83b146c..6f153db 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -3607,4 +3607,63 @@ class UnifiedMathVisitor(MathExprVisitor): def visitErfinvFunc(self, ctx): """erfinv(x) - inverse error function""" x = self._promote_to_tensor((yield ctx.expr())) - return torch.erfinv(x) \ No newline at end of file + return torch.erfinv(x) + + def visitWhereFunc(self, ctx): + cond = (yield ctx.expr(0)) + a = (yield ctx.expr(1)) + b = (yield ctx.expr(2)) + + if self._is_tensor(cond) or self._is_tensor(a) or self._is_tensor(b): + cond_t = self._promote_to_tensor(cond) + if cond_t.dtype != torch.bool: + cond_t = cond_t != 0 + a_t = self._promote_to_tensor(a) + b_t = self._promote_to_tensor(b) + return torch.where(cond_t, a_t, b_t).contiguous() + + def rec(c, av, bv): + if self._is_list(c): + out = [] + for i, ci in enumerate(c): + ai = av[i] if self._is_list(av) and i < len(av) else av + bi = bv[i] if self._is_list(bv) and i < len(bv) else bv + out.append(rec(ci, ai, bi)) + return out + return av if bool(c) else bv + + return rec(cond, a, b) + + def visitHistogramFunc(self, ctx): + x = self._promote_to_tensor((yield ctx.expr(0))).float().flatten() + bins_raw = (yield ctx.expr(1)) + min_raw = (yield ctx.expr(2)) + max_raw = (yield ctx.expr(3)) + + bins = int(bins_raw.item()) if self._is_tensor(bins_raw) else int(bins_raw) + min_v = float(min_raw.item()) if self._is_tensor(min_raw) else float(min_raw) + max_v = float(max_raw.item()) if self._is_tensor(max_raw) else float(max_raw) + + if bins <= 0: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: histogram bins must be > 0") + if max_v <= min_v: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: histogram requires max > min") + + return torch.histc(x, bins=bins, min=min_v, max=max_v).contiguous() + + def visitFlowMagFunc(self, ctx): + flow = self._promote_to_tensor((yield ctx.expr())) + if flow.shape[-1] != 2: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: flow_mag expects [..., 2], got {tuple(flow.shape)}") + dx = flow[..., 0] + dy = flow[..., 1] + return torch.sqrt(dx * dx + dy * dy).contiguous() + + def visitFlowAngFunc(self, ctx): + flow = self._promote_to_tensor((yield ctx.expr())) + if flow.shape[-1] != 2: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: flow_ang expects [..., 2], got {tuple(flow.shape)}") + dx = flow[..., 0] + dy = flow[..., 1] + + return torch.atan2(dy, dx).contiguous() \ No newline at end of file diff --git a/web/script_text_input.js b/web/script_text_input.js index 69d32f2..2067842 100644 --- a/web/script_text_input.js +++ b/web/script_text_input.js @@ -28,7 +28,7 @@ const FUNCTIONS = new Set([ "worley", "cellular_noise", "voronoi_noise", "plasma", "turbulence", "plasma_noise", "upper", "lower", "split", "join", "substring", "substr", "find", "trim", "replace", "dilate", "erode", "morph_open", "morph_close", - "rgb_to_hsv", "hsv_to_rgb", "int_to_rgb", "rgb_to_int" + "rgb_to_hsv", "hsv_to_rgb", "int_to_rgb", "rgb_to_int", "where", "histogram", "hist", "flow_mag", "flow_magnitude", "flow_ang", "flow_angle", ]); const BRACKET_PAIRS = {