add flow magnitude and angle, histogram
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
@@ -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';
|
||||
|
||||
@@ -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()
|
||||
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user