add flow magnitude and angle, histogram

This commit is contained in:
mcDandy
2026-04-20 14:24:37 +02:00
parent b8bf4a8ea7
commit 36fcbd48c8
4 changed files with 79 additions and 7 deletions
+6 -2
View File
@@ -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)
+12 -3
View File
@@ -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';
+60 -1
View File
@@ -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)
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()
+1 -1
View File
@@ -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 = {