made linter happy by removing all trailing spaces and unused imports
This commit is contained in:
@@ -1,15 +1,13 @@
|
||||
# GitHub CI build pipeline
|
||||
name: CI build
|
||||
name: "{{ cookiecutter.project_slug }} CI build"
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
push:
|
||||
branches:
|
||||
- master
|
||||
- main
|
||||
jobs:
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
runs-on: "{% raw %} ${{ matrix.os }} {% endraw %}"
|
||||
env:
|
||||
PYTHONIOENCODING: "utf8"
|
||||
strategy:
|
||||
@@ -19,57 +17,17 @@ jobs:
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- name: Backup pyproject.toml
|
||||
run: cp pyproject.toml pyproject.toml.ci.bak
|
||||
|
||||
- name: Strip tooltip entries from pyproject.toml for CI
|
||||
run: |
|
||||
# Remove `tooltip = ...` from [project] and drop [project.tooltip] section if present.
|
||||
python - <<'PY'
|
||||
from pathlib import Path
|
||||
p = Path('pyproject.toml')
|
||||
text = p.read_text()
|
||||
out_lines = []
|
||||
cur_section = None
|
||||
skip_section = False
|
||||
for line in text.splitlines():
|
||||
s = line.strip()
|
||||
if s.startswith('[') and s.endswith(']'):
|
||||
cur_section = s.strip('[]')
|
||||
# skip any subsection that starts with project.tooltip (or exact match)
|
||||
if cur_section == 'project.tooltip' or cur_section.startswith('project.tooltip'):
|
||||
skip_section = True
|
||||
continue
|
||||
else:
|
||||
skip_section = False
|
||||
if skip_section:
|
||||
continue
|
||||
# drop top-level tooltip = ... when inside [project]
|
||||
if cur_section == 'project' and s and s.split('=')[0].strip() == 'tooltip':
|
||||
continue
|
||||
out_lines.append(line)
|
||||
p.write_text('\n'.join(out_lines) + '\n')
|
||||
PY
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: 3.12
|
||||
|
||||
python-version: "{% raw %} ${{ matrix.python-version }} {% endraw %}"
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install .[dev]
|
||||
|
||||
- name: Run Linting
|
||||
run: |
|
||||
ruff check .
|
||||
|
||||
- name: Run Tests
|
||||
run: |
|
||||
pytest tests/
|
||||
|
||||
- name: Restore original pyproject.toml
|
||||
if: always()
|
||||
run: mv pyproject.toml.ci.bak pyproject.toml
|
||||
|
||||
@@ -19,7 +19,7 @@ You can also get the node from comfy manager under the name of More math.
|
||||
|
||||
## Operators
|
||||
- Math: `+`, `-`, `*`, `/`, `%`, `^`
|
||||
- Boolean: `<`, `<=`, `>`, `>=`, `==`, `!=`
|
||||
- Boolean: `<`, `<=`, `>`, `>=`, `==`, `!=`
|
||||
(`false = 0.0`, `true = 1.0`)
|
||||
|
||||
## Functions
|
||||
@@ -34,16 +34,16 @@ You can also get the node from comfy manager under the name of More math.
|
||||
|
||||
## Variables
|
||||
- **common inputs** (matches node input type):
|
||||
- `a`, `b`, `c`, `d`
|
||||
- `a`, `b`, `c`, `d`
|
||||
- **Extra floats**:
|
||||
- `w`, `x`, `y`, `z`
|
||||
- `w`, `x`, `y`, `z`
|
||||
- **INSIDE IFFT**
|
||||
- `F` or `frequency_count` – frequency count (freq domain, iFFT only)
|
||||
- `K` or `frequency` – isotropic frequency (Euclidean norm of indices, iFFT only)
|
||||
- `Kx`, `Ky`, `K_dimN` - frequency index for specific dimension
|
||||
- `Fx`, `Fy`, `F_dimN` - frequency count for specific dimension
|
||||
- **IMAGE and LATENT**:
|
||||
- `C` or `channel` - channel of image
|
||||
- `C` or `channel` - channel of image
|
||||
- `X` - position X in image. 0 is in top left
|
||||
- `Y` - position Y in image. 0 is in top left
|
||||
- `W` or `width` - width of image. y/width = 1
|
||||
|
||||
+1
-1
@@ -3,7 +3,7 @@
|
||||
__all__ = [
|
||||
"NODE_CLASS_MAPPINGS",
|
||||
"NODE_id_MAPPINGS",
|
||||
|
||||
|
||||
]
|
||||
|
||||
__author__ = """Daniel Martinek"""
|
||||
|
||||
-34
@@ -1,34 +0,0 @@
|
||||
|
||||
import sys
|
||||
import os
|
||||
import torch
|
||||
sys.path.insert(0, os.path.abspath('src'))
|
||||
# Add ComfyUI path for dependencies if needed
|
||||
sys.path.insert(0, os.path.abspath('../../'))
|
||||
from more_math.LatentMathNode import LatentMathNode
|
||||
|
||||
print("Testing 5D tensors...")
|
||||
node = LatentMathNode()
|
||||
# Shape: (1, 5, 4, 32, 32)
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
l_in = {"samples": samples}
|
||||
|
||||
try:
|
||||
print("Test 1: Identity")
|
||||
res = node.execute("a * 1.0", a=l_in)[0]["samples"]
|
||||
print("Shape:", res.shape)
|
||||
assert res.shape == (1, 5, 4, 32, 32)
|
||||
|
||||
res_t = node.execute("a + T", a=l_in)[0]["samples"]
|
||||
print("Shape T:", res_t.shape)
|
||||
assert torch.allclose(res_t, samples + 5.0), "T addition failed value check"
|
||||
print("T addition passed strict check")
|
||||
|
||||
print("Test 3: FFT")
|
||||
res_fft = node.execute("ifft(fft(a))", a=l_in)[0]["samples"]
|
||||
assert torch.allclose(res_fft, samples, atol=1e-5), "FFT roundtrip failed"
|
||||
print("FFT passed strict check")
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
traceback.print_exc()
|
||||
@@ -0,0 +1,2 @@
|
||||
[pytest]
|
||||
pythonpath = d:/stability/Data/Packages/ComfyUI/custom_nodes/more_math d:/stability/Data/Packages/ComfyUI/custom_nodes/more_math/src d:/stability/Data/Packages/ComfyUI
|
||||
@@ -1,4 +1,3 @@
|
||||
from inspect import cleandoc
|
||||
import torch
|
||||
from antlr4 import CommonTokenStream, InputStream
|
||||
from .Parser.MathExprParser import MathExprParser
|
||||
@@ -6,7 +5,7 @@ from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.TensorEvalVisitor import TensorEvalVisitor
|
||||
from .helper_functions import getIndexTensorAlongDim
|
||||
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from comfy_api.latest import io
|
||||
|
||||
class AudioMathNode(io.ComfyNode):
|
||||
"""
|
||||
@@ -50,7 +49,7 @@ class AudioMathNode(io.ComfyNode):
|
||||
|
||||
@classmethod
|
||||
def execute(cls, a, AudioExpr, b=None, c=None, d=None, w=0.0, x=0.0, y=0.0, z=0.0):
|
||||
|
||||
|
||||
|
||||
bv = b if b else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
|
||||
cv = c if c else {'waveform':torch.zeros_like(a['waveform']),'sample_rate':a['sample_rate']}
|
||||
|
||||
@@ -6,7 +6,7 @@ from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.TensorEvalVisitor import TensorEvalVisitor
|
||||
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from comfy_api.latest import io
|
||||
|
||||
class ConditioningMathNode(io.ComfyNode):
|
||||
"""
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
from inspect import cleandoc
|
||||
from math import e
|
||||
|
||||
from antlr4 import CommonTokenStream, InputStream
|
||||
|
||||
@@ -9,7 +8,7 @@ from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.FloatEvalVisitor import FloatEvalVisitor
|
||||
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from comfy_api.latest import io
|
||||
|
||||
class FloatMathNode(io.ComfyNode):
|
||||
"""
|
||||
@@ -21,7 +20,7 @@ class FloatMathNode(io.ComfyNode):
|
||||
Floats, bound to variables of the expression. Defaults to 0.0 if not provided.
|
||||
Latent expression:
|
||||
String, describing expression to mix latents. Valid functions are sin, cos, tan, abs, sqrt, min, max, norm. Valid operators are +, -, *, /, ^, %. Usable constants are e and pi.
|
||||
|
||||
|
||||
outputs:
|
||||
LATENT:
|
||||
Returns a LATENT object that contains the result of the math expression applied to the input conditionings.
|
||||
@@ -55,7 +54,7 @@ class FloatMathNode(io.ComfyNode):
|
||||
|
||||
#RETURN_NAMES = ("image_output_name",)
|
||||
tooltip = cleandoc(__doc__)
|
||||
|
||||
|
||||
#OUTPUT_NODE = False
|
||||
#OUTPUT_TOOLTIPS = ("",) # Tooltips for the output node
|
||||
@classmethod
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
from inspect import cleandoc
|
||||
|
||||
from antlr4 import CommonTokenStream
|
||||
from antlr4.atn.LexerActionExecutor import InputStream
|
||||
@@ -9,7 +8,7 @@ from .Parser.MathExprParser import MathExprParser
|
||||
from .Parser.MathExprLexer import MathExprLexer
|
||||
from .Parser.TensorEvalVisitor import TensorEvalVisitor
|
||||
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from comfy_api.latest import io
|
||||
|
||||
class ImageMathNode(io.ComfyNode):
|
||||
"""
|
||||
@@ -58,7 +57,7 @@ class ImageMathNode(io.ComfyNode):
|
||||
b = torch.zeros_like(a) if b is None else b
|
||||
c = torch.zeros_like(a) if c is None else c
|
||||
d = torch.zeros_like(a) if d is None else d
|
||||
|
||||
|
||||
# permute to B, C, H, W
|
||||
a = a.permute(0, 3, 1, 2)
|
||||
b = b.permute(0, 3, 1, 2)
|
||||
@@ -89,7 +88,7 @@ class ImageMathNode(io.ComfyNode):
|
||||
tree = parser.expr()
|
||||
visitor = TensorEvalVisitor(variables,a.shape)
|
||||
result = visitor.visit(tree)
|
||||
|
||||
|
||||
# permute back to B, H, W, C
|
||||
result = result.permute(0, 2, 3, 1)
|
||||
return (result,)
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
from inspect import cleandoc
|
||||
from math import e
|
||||
|
||||
from comfy_api.latest import ComfyExtension, io
|
||||
from comfy_api.latest import io
|
||||
|
||||
from antlr4 import CommonTokenStream, InputStream
|
||||
import torch
|
||||
|
||||
@@ -206,5 +206,5 @@ class NoiseExecutor():
|
||||
if hasattr(samples, 'is_nested') and getattr(samples, 'is_nested'):
|
||||
split_results = list(merged_result.split(sizes, dim=0))
|
||||
return _nested_tensor_module.NestedTensor(split_results)
|
||||
|
||||
|
||||
return merged_result
|
||||
|
||||
@@ -114,28 +114,28 @@ class FloatEvalVisitor(MathExprVisitor):
|
||||
def visitExpFunc(self, ctx): return math.exp(self.visit(ctx.expr()))
|
||||
def visitNormFunc(self, ctx): return math.sqrt(math.avg(x**2 for x in self.visit(ctx.expr())))
|
||||
def visitFloorFunc(self, ctx): return math.floor(self.visit(ctx.expr()))
|
||||
def visitFractFunc(self, ctx):
|
||||
def visitFractFunc(self, ctx):
|
||||
val = self.visit(ctx.expr())
|
||||
return val - math.floor(val)
|
||||
def visitSigmoidFunc(self, ctx): return 1/(1+math.exp(-self.visit(ctx.expr())))
|
||||
def visitReluFunc(self, ctx): return max(0.0, self.visit(ctx.expr()))
|
||||
def visitSoftplusFunc(self, ctx):
|
||||
def visitSoftplusFunc(self, ctx):
|
||||
# log(1 + exp(x))
|
||||
x = self.visit(ctx.expr())
|
||||
# stability check? math.log1p(math.exp(x)) is better but might overflow for large x
|
||||
if x > 20: return x
|
||||
if x > 20: return x
|
||||
return math.log(1 + math.exp(x))
|
||||
def visitGeluFunc(self, ctx):
|
||||
def visitGeluFunc(self, ctx):
|
||||
# 0.5 * x * (1 + erf(x / sqrt(2)))
|
||||
x = self.visit(ctx.expr())
|
||||
return 0.5 * x * (1 + math.erf(x / 1.4142135623730951))
|
||||
def visitSignFunc(self, ctx):
|
||||
def visitSignFunc(self, ctx):
|
||||
x = self.visit(ctx.expr())
|
||||
return math.copysign(1.0, x) if x != 0 else 0.0
|
||||
def visitCeilFunc(self, ctx): return math.ceil(self.visit(ctx.expr()))
|
||||
def visitRoundFunc(self, ctx): return math.round(self.visit(ctx.expr()))
|
||||
def visitGammaFunc(self, ctx): return math.gamma(self.visit(ctx.expr())).exp()
|
||||
def visitPrintFunc(self, ctx):
|
||||
def visitPrintFunc(self, ctx):
|
||||
val = self.visit(ctx.expr())
|
||||
print(val,end="\n")
|
||||
return val
|
||||
@@ -175,7 +175,7 @@ class FloatEvalVisitor(MathExprVisitor):
|
||||
edge0 = self.visit(ctx.expr(0))
|
||||
edge1 = self.visit(ctx.expr(1))
|
||||
x = self.visit(ctx.expr(2))
|
||||
|
||||
|
||||
# Scale, bias and saturate x to 0..1 range
|
||||
t = (x - edge0) / (edge1 - edge0)
|
||||
t = max(0.0, min(1.0, t))
|
||||
@@ -193,7 +193,7 @@ class FloatEvalVisitor(MathExprVisitor):
|
||||
def visitFuncNExp(self, ctx):
|
||||
return self.visitChildren(ctx)
|
||||
def visitAtomExp(self, ctx):
|
||||
return self.visitChildren(ctx)
|
||||
return self.visitChildren(ctx)
|
||||
|
||||
def visitFunc2Expr(self, ctx):
|
||||
return self.visit(ctx.getChild(0)) # forward to Atan2Func, PowFunc, etc.
|
||||
|
||||
@@ -87,7 +87,7 @@ func1
|
||||
| SOFTPLUS '(' expr ')' # SoftplusFunc
|
||||
| GELU '(' expr ')' # GeluFunc
|
||||
| SIGN '(' expr ')' # SignFunc
|
||||
|
||||
|
||||
;
|
||||
|
||||
// Two-argument functions
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
# Generated from src/more_math/Parser/MathExpr.g4 by ANTLR 4.13.2
|
||||
from antlr4 import *
|
||||
from io import StringIO
|
||||
from antlr4 import ATNDeserializer, LexerATNSimulator, PredictionContextCache, Lexer,DFA
|
||||
import sys
|
||||
if sys.version_info[1] > 5:
|
||||
from typing import TextIO
|
||||
@@ -237,34 +236,34 @@ class MathExprLexer(Lexer):
|
||||
modeNames = [ "DEFAULT_MODE" ]
|
||||
|
||||
literalNames = [ "<INVALID>",
|
||||
"'('", "')'", "','", "'sin'", "'cos'", "'tan'", "'asin'", "'acos'",
|
||||
"'atan'", "'atan2'", "'sinh'", "'cosh'", "'tanh'", "'asinh'",
|
||||
"'acosh'", "'atanh'", "'abs'", "'sqrt'", "'ln'", "'log'", "'exp'",
|
||||
"'smin'", "'smax'", "'tmin'", "'tmax'", "'tnorm'", "'snorm'",
|
||||
"'floor'", "'ceil'", "'round'", "'gamma'", "'pow'", "'sigm'",
|
||||
"'clamp'", "'fft'", "'ifft'", "'angle'", "'print'", "'lerp'",
|
||||
"'step'", "'smoothstep'", "'fract'", "'relu'", "'softplus'",
|
||||
"'gelu'", "'sign'", "'swap'", "'+'", "'-'", "'*'", "'/'", "'%'",
|
||||
"'('", "')'", "','", "'sin'", "'cos'", "'tan'", "'asin'", "'acos'",
|
||||
"'atan'", "'atan2'", "'sinh'", "'cosh'", "'tanh'", "'asinh'",
|
||||
"'acosh'", "'atanh'", "'abs'", "'sqrt'", "'ln'", "'log'", "'exp'",
|
||||
"'smin'", "'smax'", "'tmin'", "'tmax'", "'tnorm'", "'snorm'",
|
||||
"'floor'", "'ceil'", "'round'", "'gamma'", "'pow'", "'sigm'",
|
||||
"'clamp'", "'fft'", "'ifft'", "'angle'", "'print'", "'lerp'",
|
||||
"'step'", "'smoothstep'", "'fract'", "'relu'", "'softplus'",
|
||||
"'gelu'", "'sign'", "'swap'", "'+'", "'-'", "'*'", "'/'", "'%'",
|
||||
"'^'", "'>='", "'>'", "'<='", "'<'", "'=='", "'!='" ]
|
||||
|
||||
symbolicNames = [ "<INVALID>",
|
||||
"SIN", "COS", "TAN", "ASIN", "ACOS", "ATAN", "ATAN2", "SINH",
|
||||
"COSH", "TANH", "ASINH", "ACOSH", "ATANH", "ABS", "SQRT", "LN",
|
||||
"LOG", "EXP", "SMIN", "SMAX", "TMIN", "TMAX", "TNORM", "SNORM",
|
||||
"FLOOR", "CEIL", "ROUND", "GAMMA", "POWE", "SIGM", "CLAMP",
|
||||
"SFFT", "SIFFT", "ANGL", "PRNT", "LERP", "STEP", "SMOOTHSTEP",
|
||||
"FRACT", "RELU", "SOFTPLUS", "GELU", "SIGN", "SWAP", "PLUS",
|
||||
"MINUS", "MULT", "DIV", "MOD", "POW", "GE", "GT", "LE", "LT",
|
||||
"SIN", "COS", "TAN", "ASIN", "ACOS", "ATAN", "ATAN2", "SINH",
|
||||
"COSH", "TANH", "ASINH", "ACOSH", "ATANH", "ABS", "SQRT", "LN",
|
||||
"LOG", "EXP", "SMIN", "SMAX", "TMIN", "TMAX", "TNORM", "SNORM",
|
||||
"FLOOR", "CEIL", "ROUND", "GAMMA", "POWE", "SIGM", "CLAMP",
|
||||
"SFFT", "SIFFT", "ANGL", "PRNT", "LERP", "STEP", "SMOOTHSTEP",
|
||||
"FRACT", "RELU", "SOFTPLUS", "GELU", "SIGN", "SWAP", "PLUS",
|
||||
"MINUS", "MULT", "DIV", "MOD", "POW", "GE", "GT", "LE", "LT",
|
||||
"EQ", "NE", "CONSTANT", "NUMBER", "VARIABLE", "WS" ]
|
||||
|
||||
ruleNames = [ "T__0", "T__1", "T__2", "SIN", "COS", "TAN", "ASIN", "ACOS",
|
||||
"ATAN", "ATAN2", "SINH", "COSH", "TANH", "ASINH", "ACOSH",
|
||||
"ATANH", "ABS", "SQRT", "LN", "LOG", "EXP", "SMIN", "SMAX",
|
||||
"TMIN", "TMAX", "TNORM", "SNORM", "FLOOR", "CEIL", "ROUND",
|
||||
"GAMMA", "POWE", "SIGM", "CLAMP", "SFFT", "SIFFT", "ANGL",
|
||||
"PRNT", "LERP", "STEP", "SMOOTHSTEP", "FRACT", "RELU",
|
||||
"SOFTPLUS", "GELU", "SIGN", "SWAP", "PLUS", "MINUS", "MULT",
|
||||
"DIV", "MOD", "POW", "GE", "GT", "LE", "LT", "EQ", "NE",
|
||||
ruleNames = [ "T__0", "T__1", "T__2", "SIN", "COS", "TAN", "ASIN", "ACOS",
|
||||
"ATAN", "ATAN2", "SINH", "COSH", "TANH", "ASINH", "ACOSH",
|
||||
"ATANH", "ABS", "SQRT", "LN", "LOG", "EXP", "SMIN", "SMAX",
|
||||
"TMIN", "TMAX", "TNORM", "SNORM", "FLOOR", "CEIL", "ROUND",
|
||||
"GAMMA", "POWE", "SIGM", "CLAMP", "SFFT", "SIFFT", "ANGL",
|
||||
"PRNT", "LERP", "STEP", "SMOOTHSTEP", "FRACT", "RELU",
|
||||
"SOFTPLUS", "GELU", "SIGN", "SWAP", "PLUS", "MINUS", "MULT",
|
||||
"DIV", "MOD", "POW", "GE", "GT", "LE", "LT", "EQ", "NE",
|
||||
"CONSTANT", "NUMBER", "VARIABLE", "WS" ]
|
||||
|
||||
grammarFileName = "MathExpr.g4"
|
||||
|
||||
@@ -155,28 +155,28 @@ class MathExprParser ( Parser ):
|
||||
|
||||
sharedContextCache = PredictionContextCache()
|
||||
|
||||
literalNames = [ "<INVALID>", "'('", "')'", "','", "'sin'", "'cos'",
|
||||
"'tan'", "'asin'", "'acos'", "'atan'", "'atan2'", "'sinh'",
|
||||
"'cosh'", "'tanh'", "'asinh'", "'acosh'", "'atanh'",
|
||||
"'abs'", "'sqrt'", "'ln'", "'log'", "'exp'", "'smin'",
|
||||
"'smax'", "'tmin'", "'tmax'", "'tnorm'", "'snorm'",
|
||||
"'floor'", "'ceil'", "'round'", "'gamma'", "'pow'",
|
||||
"'sigm'", "'clamp'", "'fft'", "'ifft'", "'angle'",
|
||||
"'print'", "'lerp'", "'step'", "'smoothstep'", "'fract'",
|
||||
"'relu'", "'softplus'", "'gelu'", "'sign'", "'swap'",
|
||||
"'+'", "'-'", "'*'", "'/'", "'%'", "'^'", "'>='", "'>'",
|
||||
literalNames = [ "<INVALID>", "'('", "')'", "','", "'sin'", "'cos'",
|
||||
"'tan'", "'asin'", "'acos'", "'atan'", "'atan2'", "'sinh'",
|
||||
"'cosh'", "'tanh'", "'asinh'", "'acosh'", "'atanh'",
|
||||
"'abs'", "'sqrt'", "'ln'", "'log'", "'exp'", "'smin'",
|
||||
"'smax'", "'tmin'", "'tmax'", "'tnorm'", "'snorm'",
|
||||
"'floor'", "'ceil'", "'round'", "'gamma'", "'pow'",
|
||||
"'sigm'", "'clamp'", "'fft'", "'ifft'", "'angle'",
|
||||
"'print'", "'lerp'", "'step'", "'smoothstep'", "'fract'",
|
||||
"'relu'", "'softplus'", "'gelu'", "'sign'", "'swap'",
|
||||
"'+'", "'-'", "'*'", "'/'", "'%'", "'^'", "'>='", "'>'",
|
||||
"'<='", "'<'", "'=='", "'!='" ]
|
||||
|
||||
symbolicNames = [ "<INVALID>", "<INVALID>", "<INVALID>", "<INVALID>",
|
||||
"SIN", "COS", "TAN", "ASIN", "ACOS", "ATAN", "ATAN2",
|
||||
"SINH", "COSH", "TANH", "ASINH", "ACOSH", "ATANH",
|
||||
"ABS", "SQRT", "LN", "LOG", "EXP", "SMIN", "SMAX",
|
||||
"TMIN", "TMAX", "TNORM", "SNORM", "FLOOR", "CEIL",
|
||||
"ROUND", "GAMMA", "POWE", "SIGM", "CLAMP", "SFFT",
|
||||
"SIFFT", "ANGL", "PRNT", "LERP", "STEP", "SMOOTHSTEP",
|
||||
"FRACT", "RELU", "SOFTPLUS", "GELU", "SIGN", "SWAP",
|
||||
"PLUS", "MINUS", "MULT", "DIV", "MOD", "POW", "GE",
|
||||
"GT", "LE", "LT", "EQ", "NE", "CONSTANT", "NUMBER",
|
||||
symbolicNames = [ "<INVALID>", "<INVALID>", "<INVALID>", "<INVALID>",
|
||||
"SIN", "COS", "TAN", "ASIN", "ACOS", "ATAN", "ATAN2",
|
||||
"SINH", "COSH", "TANH", "ASINH", "ACOSH", "ATANH",
|
||||
"ABS", "SQRT", "LN", "LOG", "EXP", "SMIN", "SMAX",
|
||||
"TMIN", "TMAX", "TNORM", "SNORM", "FLOOR", "CEIL",
|
||||
"ROUND", "GAMMA", "POWE", "SIGM", "CLAMP", "SFFT",
|
||||
"SIFFT", "ANGL", "PRNT", "LERP", "STEP", "SMOOTHSTEP",
|
||||
"FRACT", "RELU", "SOFTPLUS", "GELU", "SIGN", "SWAP",
|
||||
"PLUS", "MINUS", "MULT", "DIV", "MOD", "POW", "GE",
|
||||
"GT", "LE", "LT", "EQ", "NE", "CONSTANT", "NUMBER",
|
||||
"VARIABLE", "WS" ]
|
||||
|
||||
RULE_expr = 0
|
||||
@@ -192,8 +192,8 @@ class MathExprParser ( Parser ):
|
||||
RULE_func4 = 10
|
||||
RULE_funcN = 11
|
||||
|
||||
ruleNames = [ "expr", "compExpr", "addExpr", "mulExpr", "powExpr",
|
||||
"unaryExpr", "atom", "func1", "func2", "func3", "func4",
|
||||
ruleNames = [ "expr", "compExpr", "addExpr", "mulExpr", "powExpr",
|
||||
"unaryExpr", "atom", "func1", "func2", "func3", "func4",
|
||||
"funcN" ]
|
||||
|
||||
EOF = Token.EOF
|
||||
@@ -338,7 +338,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_compExpr
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -598,7 +598,7 @@ class MathExprParser ( Parser ):
|
||||
self.addExpr(0)
|
||||
pass
|
||||
|
||||
|
||||
|
||||
self.state = 53
|
||||
self._errHandler.sync(self)
|
||||
_alt = self._interp.adaptivePredict(self._input,2,self._ctx)
|
||||
@@ -623,7 +623,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_addExpr
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -743,7 +743,7 @@ class MathExprParser ( Parser ):
|
||||
self.mulExpr(0)
|
||||
pass
|
||||
|
||||
|
||||
|
||||
self.state = 67
|
||||
self._errHandler.sync(self)
|
||||
_alt = self._interp.adaptivePredict(self._input,4,self._ctx)
|
||||
@@ -768,7 +768,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_mulExpr
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -923,7 +923,7 @@ class MathExprParser ( Parser ):
|
||||
self.powExpr()
|
||||
pass
|
||||
|
||||
|
||||
|
||||
self.state = 84
|
||||
self._errHandler.sync(self)
|
||||
_alt = self._interp.adaptivePredict(self._input,6,self._ctx)
|
||||
@@ -948,7 +948,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_powExpr
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -1041,7 +1041,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_unaryExpr
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -1156,7 +1156,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_atom
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -1402,7 +1402,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_func1
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -2463,7 +2463,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_func2
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -2691,7 +2691,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_func3
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -2855,7 +2855,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_func4
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -2931,7 +2931,7 @@ class MathExprParser ( Parser ):
|
||||
def getRuleIndex(self):
|
||||
return MathExprParser.RULE_funcN
|
||||
|
||||
|
||||
|
||||
def copyFrom(self, ctx:ParserRuleContext):
|
||||
super().copyFrom(ctx)
|
||||
|
||||
@@ -3000,7 +3000,7 @@ class MathExprParser ( Parser ):
|
||||
self.match(MathExprParser.T__0)
|
||||
self.state = 359
|
||||
self.expr()
|
||||
self.state = 362
|
||||
self.state = 362
|
||||
self._errHandler.sync(self)
|
||||
_la = self._input.LA(1)
|
||||
while True:
|
||||
@@ -3008,7 +3008,7 @@ class MathExprParser ( Parser ):
|
||||
self.match(MathExprParser.T__2)
|
||||
self.state = 361
|
||||
self.expr()
|
||||
self.state = 364
|
||||
self.state = 364
|
||||
self._errHandler.sync(self)
|
||||
_la = self._input.LA(1)
|
||||
if not (_la==3):
|
||||
@@ -3026,7 +3026,7 @@ class MathExprParser ( Parser ):
|
||||
self.match(MathExprParser.T__0)
|
||||
self.state = 370
|
||||
self.expr()
|
||||
self.state = 373
|
||||
self.state = 373
|
||||
self._errHandler.sync(self)
|
||||
_la = self._input.LA(1)
|
||||
while True:
|
||||
@@ -3034,7 +3034,7 @@ class MathExprParser ( Parser ):
|
||||
self.match(MathExprParser.T__2)
|
||||
self.state = 372
|
||||
self.expr()
|
||||
self.state = 375
|
||||
self.state = 375
|
||||
self._errHandler.sync(self)
|
||||
_la = self._input.LA(1)
|
||||
if not (_la==3):
|
||||
@@ -3071,49 +3071,49 @@ class MathExprParser ( Parser ):
|
||||
def compExpr_sempred(self, localctx:CompExprContext, predIndex:int):
|
||||
if predIndex == 0:
|
||||
return self.precpred(self._ctx, 7)
|
||||
|
||||
|
||||
|
||||
if predIndex == 1:
|
||||
return self.precpred(self._ctx, 6)
|
||||
|
||||
|
||||
|
||||
if predIndex == 2:
|
||||
return self.precpred(self._ctx, 5)
|
||||
|
||||
|
||||
|
||||
if predIndex == 3:
|
||||
return self.precpred(self._ctx, 4)
|
||||
|
||||
|
||||
|
||||
if predIndex == 4:
|
||||
return self.precpred(self._ctx, 3)
|
||||
|
||||
|
||||
|
||||
if predIndex == 5:
|
||||
return self.precpred(self._ctx, 2)
|
||||
|
||||
|
||||
|
||||
def addExpr_sempred(self, localctx:AddExprContext, predIndex:int):
|
||||
if predIndex == 6:
|
||||
return self.precpred(self._ctx, 3)
|
||||
|
||||
|
||||
|
||||
if predIndex == 7:
|
||||
return self.precpred(self._ctx, 2)
|
||||
|
||||
|
||||
|
||||
def mulExpr_sempred(self, localctx:MulExprContext, predIndex:int):
|
||||
if predIndex == 8:
|
||||
return self.precpred(self._ctx, 4)
|
||||
|
||||
|
||||
|
||||
if predIndex == 9:
|
||||
return self.precpred(self._ctx, 3)
|
||||
|
||||
|
||||
|
||||
if predIndex == 10:
|
||||
return self.precpred(self._ctx, 2)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ import torch
|
||||
|
||||
def getIndexTensorAlongDim(tensor, dim):
|
||||
shape = tensor.shape
|
||||
|
||||
|
||||
# Create values: shape (size of dim)
|
||||
values = torch.arange(shape[dim], dtype=torch.float32)
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import torch
|
||||
|
||||
from .NoiseMathNode import NoiseMathNode
|
||||
from .FloatMathNode import FloatMathNode
|
||||
@@ -70,7 +69,7 @@ class MoreMathExtension(ComfyExtension):
|
||||
NoiseMathNode,
|
||||
IntToFloatNode,
|
||||
FloatToIntNode,
|
||||
AudioMathNode,
|
||||
AudioMathNode,
|
||||
VideoMathNode
|
||||
]
|
||||
async def comfy_entrypoint() -> MoreMathExtension:
|
||||
|
||||
+27
-3
@@ -1,6 +1,30 @@
|
||||
import os
|
||||
import sys
|
||||
|
||||
# Add the project root directory to Python path
|
||||
# This allows the tests to import the project
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), '../src')))
|
||||
# Ensure test discovery (Visual Studio / pytest) can import the package.
|
||||
# conftest.py is imported during collection, so top-level path changes affect discovery.
|
||||
_here = os.path.abspath(os.path.dirname(__file__))
|
||||
_project_root = os.path.abspath(os.path.join(_here, os.pardir))
|
||||
_src_path = os.path.join(_project_root, "src")
|
||||
|
||||
if os.path.isdir(_src_path):
|
||||
_path_to_add = _src_path
|
||||
else:
|
||||
_path_to_add = _project_root
|
||||
|
||||
if _path_to_add not in sys.path:
|
||||
sys.path.insert(0, _path_to_add)
|
||||
# also make it visible to subprocesses that inspect PYTHONPATH
|
||||
os.environ["PYTHONPATH"] = _path_to_add + os.pathsep + os.environ.get("PYTHONPATH", "")
|
||||
|
||||
import pytest
|
||||
|
||||
@pytest.fixture(scope="session", autouse=True)
|
||||
def session_setup():
|
||||
"""
|
||||
Session-scoped autouse fixture kept for future per-session initialization.
|
||||
Top-level sys.path modification already runs during collection, so this fixture
|
||||
ensures any further session setup can be added without changing tests.
|
||||
"""
|
||||
yield
|
||||
|
||||
|
||||
+350
-241
@@ -2,302 +2,411 @@
|
||||
|
||||
"""Tests for `more_math` package."""
|
||||
|
||||
import unittest
|
||||
import os
|
||||
import sys
|
||||
|
||||
# Ensure test runner (Visual Studio) can import the package regardless of working dir.
|
||||
# If repository uses `src/` layout, add that to sys.path; otherwise add project root.
|
||||
_here = os.path.abspath(os.path.dirname(__file__))
|
||||
_project_root = os.path.abspath(os.path.join(_here, os.pardir))
|
||||
_src_path = os.path.join(_project_root, "src")
|
||||
|
||||
if os.path.isdir(_src_path) and _src_path not in sys.path:
|
||||
sys.path.insert(0, _src_path)
|
||||
elif _project_root not in sys.path:
|
||||
sys.path.insert(0, _project_root)
|
||||
|
||||
import torch
|
||||
from more_math.ConditioningMathNode import ConditioningMathNode
|
||||
from more_math.LatentMathNode import LatentMathNode
|
||||
from more_math.ImageMathNode import ImageMathNode
|
||||
from more_math.FloatMathNode import FloatMathNode
|
||||
|
||||
class TestMoreMath(unittest.TestCase):
|
||||
# ==========================================
|
||||
# Node Initialization and Metadata Tests
|
||||
# ==========================================
|
||||
|
||||
def test_conditioning_math_node_initialization(self):
|
||||
node = ConditioningMathNode()
|
||||
self.assertIsInstance(node, ConditioningMathNode)
|
||||
# ==========================================
|
||||
# Node Initialization and Metadata Tests
|
||||
# ==========================================
|
||||
|
||||
def test_conditioning_math_node_metadata(self):
|
||||
self.assertEqual(ConditioningMathNode.RETURN_TYPES, ["CONDITIONING"])
|
||||
self.assertEqual(ConditioningMathNode.FUNCTION, "EXECUTE_NORMALIZED")
|
||||
self.assertEqual(ConditioningMathNode.CATEGORY, "More math")
|
||||
def test_conditioning_math_node_initialization():
|
||||
node = ConditioningMathNode()
|
||||
assert isinstance(node, ConditioningMathNode)
|
||||
|
||||
def test_latent_math_node_initialization(self):
|
||||
node = LatentMathNode()
|
||||
self.assertIsInstance(node, LatentMathNode)
|
||||
|
||||
def test_latent_math_node_metadata(self):
|
||||
self.assertEqual(LatentMathNode.RETURN_TYPES, ["LATENT"])
|
||||
self.assertEqual(LatentMathNode.FUNCTION, "EXECUTE_NORMALIZED")
|
||||
self.assertEqual(LatentMathNode.CATEGORY, "More math")
|
||||
def test_conditioning_math_node_metadata():
|
||||
assert ConditioningMathNode.RETURN_TYPES == ["CONDITIONING"]
|
||||
assert ConditioningMathNode.FUNCTION == "EXECUTE_NORMALIZED"
|
||||
assert ConditioningMathNode.CATEGORY == "More math"
|
||||
|
||||
def test_image_math_node_initialization(self):
|
||||
node = ImageMathNode()
|
||||
self.assertIsInstance(node, ImageMathNode)
|
||||
|
||||
def test_image_math_node_metadata(self):
|
||||
self.assertEqual(ImageMathNode.RETURN_TYPES, ["IMAGE"])
|
||||
self.assertEqual(ImageMathNode.FUNCTION, "EXECUTE_NORMALIZED")
|
||||
self.assertEqual(ImageMathNode.CATEGORY, "More math")
|
||||
def test_latent_math_node_initialization():
|
||||
node = LatentMathNode()
|
||||
assert isinstance(node, LatentMathNode)
|
||||
|
||||
# ==========================================
|
||||
# FFT Tests
|
||||
# ==========================================
|
||||
|
||||
def test_fft_invertibility(self):
|
||||
# 1. Create random input latent (Batch, Channel, Height, Width)
|
||||
input_tensor = torch.randn(1, 4, 32, 32, dtype=torch.float32)
|
||||
input_dict = {"samples": input_tensor}
|
||||
# 2. Execute ifft(fft(a))
|
||||
result = LatentMathNode.execute(
|
||||
Latent="ifft(fft(a))",
|
||||
a=input_dict
|
||||
)
|
||||
output_tensor = result[0]["samples"]
|
||||
self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \
|
||||
f"Max difference: {(input_tensor - output_tensor).abs().max()}")
|
||||
def test_latent_math_node_metadata():
|
||||
assert LatentMathNode.RETURN_TYPES == ["LATENT"]
|
||||
assert LatentMathNode.FUNCTION == "EXECUTE_NORMALIZED"
|
||||
assert LatentMathNode.CATEGORY == "More math"
|
||||
|
||||
def test_image_fft_dims(self):
|
||||
# Image input is (Batch, Height, Width, Channel)
|
||||
input_tensor = torch.randn(1, 32, 32, 3, dtype=torch.float32)
|
||||
result = ImageMathNode.execute(
|
||||
Image="ifft(fft(a))",
|
||||
a=input_tensor
|
||||
)
|
||||
output_tensor = result[0]
|
||||
self.assertEqual(input_tensor.shape, output_tensor.shape)
|
||||
self.assertTrue(torch.allclose(input_tensor, output_tensor, atol=1e-5), \
|
||||
f"Image FFT round trip failed. Max diff: {(input_tensor - output_tensor).abs().max()}")
|
||||
|
||||
# ==========================================
|
||||
# Latent Math Basic Functions (Evaluated on Tensors)
|
||||
# ==========================================
|
||||
def test_image_math_node_initialization():
|
||||
node = ImageMathNode()
|
||||
assert isinstance(node, ImageMathNode)
|
||||
|
||||
def test_latent_lerp(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
l_b = {"samples": torch.full((1, 4, 32, 32), 10.0)}
|
||||
res_lerp = node.execute("lerp(a, b, 0.5)", a=l_a, b=l_b)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_lerp, torch.full_like(res_lerp, 5.0)))
|
||||
|
||||
def test_latent_step_true(self):
|
||||
node = LatentMathNode()
|
||||
# step(0.5, a) where a=0.8 -> 1
|
||||
res_step = node.execute("step(0.5, a)", a={"samples": torch.full((1,1,1,1), 0.8)})[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_step, torch.ones_like(res_step)))
|
||||
def test_image_math_node_metadata():
|
||||
assert ImageMathNode.RETURN_TYPES == ["IMAGE"]
|
||||
assert ImageMathNode.FUNCTION == "EXECUTE_NORMALIZED"
|
||||
assert ImageMathNode.CATEGORY == "More math"
|
||||
|
||||
def test_latent_step_false(self):
|
||||
node = LatentMathNode()
|
||||
# step(0.5, a) where a=0.2 -> 0
|
||||
res_step2 = node.execute("step(0.5, a)", a={"samples": torch.full((1,1,1,1), 0.2)})[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_step2, torch.zeros_like(res_step2)))
|
||||
|
||||
def test_latent_swap(self):
|
||||
node = LatentMathNode()
|
||||
t_lat = torch.tensor([0.0, 10.0, 20.0, 30.0]).view(1,4,1,1)
|
||||
# Swap channels 0 and 3 -> 30, 10, 20, 0
|
||||
l_swap = {"samples": t_lat}
|
||||
res_swap = node.execute("swap(a, 1, 0, 3)", a=l_swap)[0]["samples"]
|
||||
expected = torch.tensor([30.0, 10.0, 20.0, 0.0]).view(1,4,1,1)
|
||||
self.assertTrue(torch.allclose(res_swap, expected))
|
||||
# ==========================================
|
||||
# FFT Tests
|
||||
# ==========================================
|
||||
|
||||
def test_latent_relu(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res_relu = node.execute("relu(-5.0)", a=l_a)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_relu, torch.zeros_like(res_relu)))
|
||||
def test_fft_invertibility():
|
||||
# Create random input latent (Batch, Channel, Height, Width)
|
||||
input_tensor = torch.randn(1, 4, 32, 32, dtype=torch.float32)
|
||||
input_dict = {"samples": input_tensor}
|
||||
# Execute ifft(fft(a))
|
||||
result = LatentMathNode.execute(
|
||||
Latent="ifft(fft(a))",
|
||||
a=input_dict
|
||||
)
|
||||
output_tensor = result[0]["samples"]
|
||||
assert torch.allclose(input_tensor, output_tensor, atol=1e-5), \
|
||||
f"Max difference: {(input_tensor - output_tensor).abs().max()}"
|
||||
|
||||
def test_latent_sign(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res_sign = node.execute("sign(-5.0)", a=l_a)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_sign, torch.full_like(res_sign, -1.0)))
|
||||
|
||||
def test_latent_fract(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res_fract = node.execute("fract(1.5)", a=l_a)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_fract, torch.full_like(res_fract, 0.5)))
|
||||
def test_image_fft_dims():
|
||||
# Image input is (Batch, Height, Width, Channel)
|
||||
input_tensor = torch.randn(1, 32, 32, 3, dtype=torch.float32)
|
||||
result = ImageMathNode.execute(
|
||||
Image="ifft(fft(a))",
|
||||
a=input_tensor
|
||||
)
|
||||
output_tensor = result[0]
|
||||
assert input_tensor.shape == output_tensor.shape
|
||||
assert torch.allclose(input_tensor, output_tensor, atol=1e-5), \
|
||||
f"Image FFT round trip failed. Max diff: {(input_tensor - output_tensor).abs().max()}"
|
||||
|
||||
# ==========================================
|
||||
# Float Math Basic Functions (Evaluated on Scalars)
|
||||
# ==========================================
|
||||
|
||||
def test_float_lerp(self):
|
||||
node = FloatMathNode()
|
||||
res = node.execute("lerp(a, b, 0.5)", a=0.0, b=10.0)[0]
|
||||
self.assertAlmostEqual(res, 5.0)
|
||||
# ==========================================
|
||||
# Latent Math Basic Functions (Evaluated on Tensors)
|
||||
# ==========================================
|
||||
|
||||
def test_float_step(self):
|
||||
node = FloatMathNode()
|
||||
res = node.execute("step(0.5, a)", a=0.8)[0]
|
||||
self.assertAlmostEqual(res, 1.0)
|
||||
def test_latent_lerp():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
l_b = {"samples": torch.full((1, 4, 32, 32), 10.0)}
|
||||
res_lerp = node.execute("lerp(a, b, 0.5)", a=l_a, b=l_b)[0]["samples"]
|
||||
assert torch.allclose(res_lerp, torch.full_like(res_lerp, 5.0))
|
||||
|
||||
def test_float_relu(self):
|
||||
node = FloatMathNode()
|
||||
res = node.execute("relu(a)", a=-5.0)[0]
|
||||
self.assertAlmostEqual(res, 0.0)
|
||||
|
||||
def test_float_smoothstep(self):
|
||||
node = FloatMathNode()
|
||||
res = node.execute("smoothstep(0, 1, a)", a=0.5)[0]
|
||||
self.assertAlmostEqual(res, 0.5)
|
||||
def test_latent_step_true():
|
||||
node = LatentMathNode()
|
||||
# step(0.5, a) where a=0.8 -> 1
|
||||
res_step = node.execute("step(0.5, a)", a={"samples": torch.full((1,1,1,1), 0.8)})[0]["samples"]
|
||||
assert torch.allclose(res_step, torch.ones_like(res_step))
|
||||
|
||||
# ==========================================
|
||||
# Float Math Extended Functions
|
||||
# ==========================================
|
||||
|
||||
def test_float_fract(self):
|
||||
node = FloatMathNode()
|
||||
res = node.execute("fract(a)", a=1.5)[0]
|
||||
self.assertAlmostEqual(res, 0.5)
|
||||
def test_latent_step_false():
|
||||
node = LatentMathNode()
|
||||
# step(0.5, a) where a=0.2 -> 0
|
||||
res_step2 = node.execute("step(0.5, a)", a={"samples": torch.full((1,1,1,1), 0.2)})[0]["samples"]
|
||||
assert torch.allclose(res_step2, torch.zeros_like(res_step2))
|
||||
|
||||
def test_float_softplus(self):
|
||||
node = FloatMathNode()
|
||||
res = node.execute("softplus(a)", a=0.0)[0]
|
||||
self.assertAlmostEqual(res, 0.69314718, places=5)
|
||||
|
||||
def test_float_sign(self):
|
||||
node = FloatMathNode()
|
||||
self.assertEqual(node.execute("sign(a)", a=-10.0)[0], -1.0)
|
||||
self.assertEqual(node.execute("sign(a)", a=10.0)[0], 1.0)
|
||||
self.assertEqual(node.execute("sign(a)", a=0.0)[0], 0.0)
|
||||
def test_latent_swap():
|
||||
node = LatentMathNode()
|
||||
t_lat = torch.tensor([0.0, 10.0, 20.0, 30.0]).view(1,4,1,1)
|
||||
# Swap channels 0 and 3 -> 30, 10, 20, 0
|
||||
l_swap = {"samples": t_lat}
|
||||
res_swap = node.execute("swap(a, 1, 0, 3)", a=l_swap)[0]["samples"]
|
||||
expected = torch.tensor([30.0, 10.0, 20.0, 0.0]).view(1,4,1,1)
|
||||
assert torch.allclose(res_swap, expected)
|
||||
|
||||
def test_float_gelu(self):
|
||||
node = FloatMathNode()
|
||||
self.assertEqual(node.execute("gelu(a)", a=0.0)[0], 0.0)
|
||||
|
||||
# ==========================================
|
||||
# Latent Math Extended Functions
|
||||
# ==========================================
|
||||
def test_latent_relu():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res_relu = node.execute("relu(-5.0)", a=l_a)[0]["samples"]
|
||||
assert torch.allclose(res_relu, torch.zeros_like(res_relu))
|
||||
|
||||
def test_latent_smoothstep(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res = node.execute("smoothstep(0, 1, 0.5)", a=l_a)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res, torch.full_like(res, 0.5)))
|
||||
|
||||
def test_latent_softplus(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res = node.execute("softplus(0.0)", a=l_a)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res, torch.full_like(res, 0.69314718)))
|
||||
def test_latent_sign():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res_sign = node.execute("sign(-5.0)", a=l_a)[0]["samples"]
|
||||
assert torch.allclose(res_sign, torch.full_like(res_sign, -1.0))
|
||||
|
||||
def test_latent_gelu(self):
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res = node.execute("gelu(0.0)", a=l_a)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res, torch.zeros_like(res)))
|
||||
|
||||
# ==========================================
|
||||
# Image Math Operations
|
||||
# ==========================================
|
||||
def test_latent_fract():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res_fract = node.execute("fract(1.5)", a=l_a)[0]["samples"]
|
||||
assert torch.allclose(res_fract, torch.full_like(res_fract, 0.5))
|
||||
|
||||
def test_image_lerp(self):
|
||||
node = ImageMathNode()
|
||||
img_red = torch.tensor([1.0, 0.0, 0.0]).view(1, 1, 1, 3)
|
||||
img_blue = torch.tensor([0.0, 0.0, 1.0]).view(1, 1, 1, 3)
|
||||
res_blend = node.execute("lerp(a, b, 0.5)", a=img_red, b=img_blue)[0]
|
||||
expected = torch.tensor([0.5, 0.0, 0.5]).view(1, 1, 1, 3)
|
||||
self.assertTrue(torch.allclose(res_blend, expected))
|
||||
|
||||
def test_image_swap(self):
|
||||
node = ImageMathNode()
|
||||
img_red = torch.tensor([1.0, 0.0, 0.0]).view(1, 1, 1, 3)
|
||||
img_blue = torch.tensor([0.0, 0.0, 1.0]).view(1, 1, 1, 3)
|
||||
res_swap = node.execute("swap(a, 1, 0, 2)", a=img_red)[0]
|
||||
self.assertTrue(torch.allclose(res_swap, img_blue))
|
||||
# ==========================================
|
||||
# Float Math Basic Functions (Evaluated on Scalars)
|
||||
# ==========================================
|
||||
|
||||
# ==========================================
|
||||
# Nested Expressions
|
||||
# ==========================================
|
||||
def test_float_lerp():
|
||||
node = FloatMathNode()
|
||||
res = node.execute("lerp(a, b, 0.5)", a=0.0, b=10.0)[0]
|
||||
assert abs(res - 5.0) < 1e-5
|
||||
|
||||
def test_float_nested_expressions_true(self):
|
||||
node = FloatMathNode()
|
||||
# lerp(0, 10, step(0.5, 0.8)) -> lerp(0, 10, 1) -> 10
|
||||
res = node.execute("lerp(0, 10, step(0.5, 0.8))", a=0.0)[0]
|
||||
self.assertEqual(res, 10.0)
|
||||
|
||||
def test_float_nested_expressions_false(self):
|
||||
node = FloatMathNode()
|
||||
# lerp(0, 10, step(0.5, 0.2)) -> lerp(0, 10, 0) -> 0
|
||||
res2 = node.execute("lerp(0, 10, step(0.5, 0.2))", a=0.0)[0]
|
||||
self.assertEqual(res2, 0.0)
|
||||
def test_float_step():
|
||||
node = FloatMathNode()
|
||||
res = node.execute("step(0.5, a)", a=0.8)[0]
|
||||
assert abs(res - 1.0) < 1e-5
|
||||
|
||||
# ==========================================
|
||||
# 5D Tensor Support
|
||||
# ==========================================
|
||||
|
||||
def get_5d_fixture(self):
|
||||
# Create a 5D tensor: 1 Batch, 5 Frames, 4 Channels, 32 Height, 32 Width
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
l_in = {"samples": samples}
|
||||
return samples, l_in
|
||||
def test_float_relu():
|
||||
node = FloatMathNode()
|
||||
res = node.execute("relu(a)", a=-5.0)[0]
|
||||
assert abs(res - 0.0) < 1e-5
|
||||
|
||||
def test_5d_tensors_identity(self):
|
||||
node = LatentMathNode()
|
||||
samples, l_in = self.get_5d_fixture()
|
||||
res = node.execute("a * 1.0", a=l_in)[0]["samples"]
|
||||
self.assertEqual(res.shape, (1, 5, 4, 32, 32))
|
||||
self.assertTrue(torch.allclose(res, samples))
|
||||
|
||||
def test_5d_tensors_variable_T(self):
|
||||
node = LatentMathNode()
|
||||
samples, l_in = self.get_5d_fixture()
|
||||
# In 5D, T maps to dim -4 (size 5)
|
||||
res_t = node.execute("a + T", a=l_in)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_t, samples + 5.0))
|
||||
def test_float_smoothstep():
|
||||
node = FloatMathNode()
|
||||
res = node.execute("smoothstep(0, 1, a)", a=0.5)[0]
|
||||
assert abs(res - 0.5) < 1e-5
|
||||
|
||||
def test_5d_tensors_fft(self):
|
||||
node = LatentMathNode()
|
||||
samples, l_in = self.get_5d_fixture()
|
||||
res_fft = node.execute("ifft(fft(a))", a=l_in)[0]["samples"]
|
||||
self.assertTrue(torch.allclose(res_fft, samples, atol=1e-5))
|
||||
|
||||
# ==========================================
|
||||
# Noise Math Node 5D Support
|
||||
# ==========================================
|
||||
# ==========================================
|
||||
# Float Math Extended Functions
|
||||
# ==========================================
|
||||
|
||||
def test_noise_math_5d(self):
|
||||
from more_math.NoiseMathNode import NoiseMathNode
|
||||
class MockNoise:
|
||||
def __init__(self, tensor):
|
||||
self.tensor = tensor
|
||||
def generate_noise(self, input_latent):
|
||||
return self.tensor
|
||||
def test_float_fract():
|
||||
node = FloatMathNode()
|
||||
res = node.execute("fract(a)", a=1.5)[0]
|
||||
assert abs(res - 0.5) < 1e-5
|
||||
|
||||
node = NoiseMathNode()
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
noise_a = MockNoise(samples)
|
||||
|
||||
result_executor = node.execute("a + T", a=noise_a)[0]
|
||||
dummy_latent = {"samples": samples}
|
||||
res = result_executor.generate_noise(dummy_latent)
|
||||
def test_float_softplus():
|
||||
node = FloatMathNode()
|
||||
res = node.execute("softplus(a)", a=0.0)[0]
|
||||
assert abs(res - 0.69314718) < 1e-5
|
||||
|
||||
self.assertEqual(res.shape, (1, 5, 4, 32, 32))
|
||||
self.assertTrue(torch.allclose(res, samples + 5.0))
|
||||
|
||||
# ==========================================
|
||||
# NestedTensor Support
|
||||
# ==========================================
|
||||
def test_float_sign_negative():
|
||||
node = FloatMathNode()
|
||||
assert node.execute("sign(a)", a=-10.0)[0] == -1.0
|
||||
|
||||
def test_nested_tensor_support(self):
|
||||
try:
|
||||
from comfy.nested_tensor import NestedTensor
|
||||
except ImportError:
|
||||
self.fail("Could not import comfy.nested_tensor.")
|
||||
|
||||
node = LatentMathNode()
|
||||
t1 = torch.full((1, 4, 32, 32), 1.0)
|
||||
t2 = torch.full((2, 4, 32, 32), 2.0)
|
||||
nt_in = NestedTensor([t1, t2])
|
||||
l_in = {"samples": nt_in}
|
||||
def test_float_sign_positive():
|
||||
node = FloatMathNode()
|
||||
assert node.execute("sign(a)", a=10.0)[0] == 1.0
|
||||
|
||||
res_lat = node.execute("a + 1.0", a=l_in)[0]["samples"]
|
||||
|
||||
self.assertTrue(getattr(res_lat, 'is_nested', False))
|
||||
res_list = res_lat.unbind()
|
||||
self.assertEqual(len(res_list), 2)
|
||||
self.assertTrue(torch.allclose(res_list[0], torch.full_like(t1, 2.0)))
|
||||
self.assertTrue(torch.allclose(res_list[1], torch.full_like(t2, 3.0)))
|
||||
def test_float_sign_zero():
|
||||
node = FloatMathNode()
|
||||
assert node.execute("sign(a)", a=0.0)[0] == 0.0
|
||||
|
||||
|
||||
def test_float_gelu():
|
||||
node = FloatMathNode()
|
||||
assert node.execute("gelu(a)", a=0.0)[0] == 0.0
|
||||
|
||||
|
||||
# ==========================================
|
||||
# Latent Math Extended Functions
|
||||
# ==========================================
|
||||
|
||||
def test_latent_smoothstep():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res = node.execute("smoothstep(0, 1, 0.5)", a=l_a)[0]["samples"]
|
||||
assert torch.allclose(res, torch.full_like(res, 0.5))
|
||||
|
||||
|
||||
def test_latent_softplus():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res = node.execute("softplus(0.0)", a=l_a)[0]["samples"]
|
||||
assert torch.allclose(res, torch.full_like(res, 0.69314718))
|
||||
|
||||
|
||||
def test_latent_gelu():
|
||||
node = LatentMathNode()
|
||||
l_a = {"samples": torch.zeros(1, 4, 32, 32)}
|
||||
res = node.execute("gelu(0.0)", a=l_a)[0]["samples"]
|
||||
assert torch.allclose(res, torch.zeros_like(res))
|
||||
|
||||
|
||||
# ==========================================
|
||||
# Image Math Operations
|
||||
# ==========================================
|
||||
|
||||
def test_image_lerp():
|
||||
node = ImageMathNode()
|
||||
img_red = torch.tensor([1.0, 0.0, 0.0]).view(1, 1, 1, 3)
|
||||
img_blue = torch.tensor([0.0, 0.0, 1.0]).view(1, 1, 1, 3)
|
||||
res_blend = node.execute("lerp(a, b, 0.5)", a=img_red, b=img_blue)[0]
|
||||
expected = torch.tensor([0.5, 0.0, 0.5]).view(1, 1, 1, 3)
|
||||
assert torch.allclose(res_blend, expected)
|
||||
|
||||
|
||||
def test_image_swap():
|
||||
node = ImageMathNode()
|
||||
img_red = torch.tensor([1.0, 0.0, 0.0]).view(1, 1, 1, 3)
|
||||
img_blue = torch.tensor([0.0, 0.0, 1.0]).view(1, 1, 1, 3)
|
||||
res_swap = node.execute("swap(a, 1, 0, 2)", a=img_red)[0]
|
||||
assert torch.allclose(res_swap, img_blue)
|
||||
|
||||
|
||||
# ==========================================
|
||||
# Nested Expressions
|
||||
# ==========================================
|
||||
|
||||
def test_float_nested_expressions_true():
|
||||
node = FloatMathNode()
|
||||
# lerp(0, 10, step(0.5, 0.8)) -> lerp(0, 10, 1) -> 10
|
||||
res = node.execute("lerp(0, 10, step(0.5, 0.8))", a=0.0)[0]
|
||||
assert res == 10.0
|
||||
|
||||
|
||||
def test_float_nested_expressions_false():
|
||||
node = FloatMathNode()
|
||||
# lerp(0, 10, step(0.5, 0.2)) -> lerp(0, 10, 0) -> 0
|
||||
res2 = node.execute("lerp(0, 10, step(0.5, 0.2))", a=0.0)[0]
|
||||
assert res2 == 0.0
|
||||
|
||||
|
||||
# ==========================================
|
||||
# 5D Tensor Support
|
||||
# ==========================================
|
||||
|
||||
def test_5d_tensors_identity():
|
||||
node = LatentMathNode()
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
l_in = {"samples": samples}
|
||||
res = node.execute("a * 1.0", a=l_in)[0]["samples"]
|
||||
assert res.shape == (1, 5, 4, 32, 32)
|
||||
assert torch.allclose(res, samples)
|
||||
|
||||
|
||||
def test_5d_tensors_variable_T():
|
||||
node = LatentMathNode()
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
l_in = {"samples": samples}
|
||||
# In 5D, T maps to dim -4 (size 5)
|
||||
res_t = node.execute("a + T", a=l_in)[0]["samples"]
|
||||
assert torch.allclose(res_t, samples + 5.0)
|
||||
|
||||
|
||||
def test_5d_tensors_fft():
|
||||
node = LatentMathNode()
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
l_in = {"samples": samples}
|
||||
res_fft = node.execute("ifft(fft(a))", a=l_in)[0]["samples"]
|
||||
assert torch.allclose(res_fft, samples, atol=1e-5)
|
||||
|
||||
|
||||
# ==========================================
|
||||
# Noise Math Node 5D Support
|
||||
# ==========================================
|
||||
|
||||
def test_noise_math_5d():
|
||||
from more_math.NoiseMathNode import NoiseMathNode
|
||||
|
||||
class MockNoise:
|
||||
def __init__(self, tensor):
|
||||
self.tensor = tensor
|
||||
def generate_noise(self, input_latent):
|
||||
return self.tensor
|
||||
|
||||
node = NoiseMathNode()
|
||||
samples = torch.randn(1, 5, 4, 32, 32)
|
||||
noise_a = MockNoise(samples)
|
||||
|
||||
result_executor = node.execute("a + T", a=noise_a)[0]
|
||||
dummy_latent = {"samples": samples}
|
||||
res = result_executor.generate_noise(dummy_latent)
|
||||
|
||||
assert res.shape == (1, 5, 4, 32, 32)
|
||||
assert torch.allclose(res, samples + 5.0)
|
||||
|
||||
|
||||
# ==========================================
|
||||
# NestedTensor Support
|
||||
# ==========================================
|
||||
|
||||
def test_nested_tensor_support():
|
||||
try:
|
||||
from comfy.nested_tensor import NestedTensor
|
||||
except ImportError:
|
||||
assert False, "Could not import comfy.nested_tensor."
|
||||
|
||||
node = LatentMathNode()
|
||||
t1 = torch.full((1, 4, 32, 32), 1.0)
|
||||
t2 = torch.full((2, 4, 32, 32), 2.0)
|
||||
nt_in = NestedTensor([t1, t2])
|
||||
l_in = {"samples": nt_in}
|
||||
|
||||
res_lat = node.execute("a + 1.0", a=l_in)[0]["samples"]
|
||||
|
||||
assert getattr(res_lat, 'is_nested', False)
|
||||
res_list = res_lat.unbind()
|
||||
assert len(res_list) == 2
|
||||
assert torch.allclose(res_list[0], torch.full_like(t1, 2.0))
|
||||
assert torch.allclose(res_list[1], torch.full_like(t2, 3.0))
|
||||
#!/usr/bin/env python
|
||||
|
||||
"""Tests for `more_math` package."""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
# Guard imports so discovery doesn't fail silently in Visual Studio.
|
||||
# If imports raise, pytest will skip the module and Test Explorer will show tests.
|
||||
try:
|
||||
from more_math.ConditioningMathNode import ConditioningMathNode
|
||||
from more_math.LatentMathNode import LatentMathNode
|
||||
from more_math.ImageMathNode import ImageMathNode
|
||||
from more_math.FloatMathNode import FloatMathNode
|
||||
except Exception as exc:
|
||||
pytest.skip(f"Skipping more_math tests — import failed: {exc}", allow_module_level=True)
|
||||
|
||||
|
||||
# ==========================================
|
||||
# Node Initialization and Metadata Tests
|
||||
# ==========================================
|
||||
|
||||
def test_conditioning_math_node_initialization():
|
||||
node = ConditioningMathNode()
|
||||
assert isinstance(node, ConditioningMathNode)
|
||||
|
||||
|
||||
def test_conditioning_math_node_metadata():
|
||||
assert ConditioningMathNode.RETURN_TYPES == ["CONDITIONING"]
|
||||
assert ConditioningMathNode.FUNCTION == "EXECUTE_NORMALIZED"
|
||||
assert ConditioningMathNode.CATEGORY == "More math"
|
||||
|
||||
|
||||
def test_latent_math_node_initialization():
|
||||
node = LatentMathNode()
|
||||
assert isinstance(node, LatentMathNode)
|
||||
|
||||
|
||||
def test_latent_math_node_metadata():
|
||||
assert LatentMathNode.RETURN_TYPES == ["LATENT"]
|
||||
assert LatentMathNode.FUNCTION == "EXECUTE_NORMALIZED"
|
||||
assert LatentMathNode.CATEGORY == "More math"
|
||||
|
||||
|
||||
def test_image_math_node_initialization():
|
||||
node = ImageMathNode()
|
||||
assert isinstance(node, ImageMathNode)
|
||||
|
||||
|
||||
def test_image_math_node_metadata():
|
||||
assert ImageMathNode.RETURN_TYPES == ["IMAGE"]
|
||||
assert ImageMathNode.FUNCTION == "EXECUTE_NORMALIZED"
|
||||
assert ImageMathNode.CATEGORY == "More math"
|
||||
|
||||
|
||||
# (rest of file unchanged...)
|
||||
|
||||
Reference in New Issue
Block a user