diff --git a/.github/workflows/build-pipeline.yml b/.github/workflows/build-pipeline.yml index 39ff3bc..52dba66 100644 --- a/.github/workflows/build-pipeline.yml +++ b/.github/workflows/build-pipeline.yml @@ -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 diff --git a/README.md b/README.md index a49ac5d..f876f07 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/__init__.py b/__init__.py index 5e16eb4..e6edefc 100644 --- a/__init__.py +++ b/__init__.py @@ -3,7 +3,7 @@ __all__ = [ "NODE_CLASS_MAPPINGS", "NODE_id_MAPPINGS", - + ] __author__ = """Daniel Martinek""" diff --git a/debug_5d.py b/debug_5d.py deleted file mode 100644 index 21452f3..0000000 --- a/debug_5d.py +++ /dev/null @@ -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() diff --git a/pytest.ini b/pytest.ini new file mode 100644 index 0000000..e4b09bf --- /dev/null +++ b/pytest.ini @@ -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 diff --git a/src/more_math/AudioMathNode.py b/src/more_math/AudioMathNode.py index 6fb432e..85835fe 100644 --- a/src/more_math/AudioMathNode.py +++ b/src/more_math/AudioMathNode.py @@ -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']} diff --git a/src/more_math/ConditioningMathNode.py b/src/more_math/ConditioningMathNode.py index 01b00f2..2d6f67e 100644 --- a/src/more_math/ConditioningMathNode.py +++ b/src/more_math/ConditioningMathNode.py @@ -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): """ diff --git a/src/more_math/FloatMathNode.py b/src/more_math/FloatMathNode.py index 5b23282..2a04da9 100644 --- a/src/more_math/FloatMathNode.py +++ b/src/more_math/FloatMathNode.py @@ -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 diff --git a/src/more_math/ImageMathNode.py b/src/more_math/ImageMathNode.py index 671e81c..b75eced 100644 --- a/src/more_math/ImageMathNode.py +++ b/src/more_math/ImageMathNode.py @@ -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,) diff --git a/src/more_math/LatentMathNode.py b/src/more_math/LatentMathNode.py index 5c10166..588a688 100644 --- a/src/more_math/LatentMathNode.py +++ b/src/more_math/LatentMathNode.py @@ -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 diff --git a/src/more_math/NoiseMathNode.py b/src/more_math/NoiseMathNode.py index b3e18e5..9e4f4c1 100644 --- a/src/more_math/NoiseMathNode.py +++ b/src/more_math/NoiseMathNode.py @@ -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 diff --git a/src/more_math/Parser/FloatEvalVisitor.py b/src/more_math/Parser/FloatEvalVisitor.py index 1a318f3..e76a464 100644 --- a/src/more_math/Parser/FloatEvalVisitor.py +++ b/src/more_math/Parser/FloatEvalVisitor.py @@ -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. diff --git a/src/more_math/Parser/MathExpr.g4 b/src/more_math/Parser/MathExpr.g4 index 13c13b6..5151296 100644 --- a/src/more_math/Parser/MathExpr.g4 +++ b/src/more_math/Parser/MathExpr.g4 @@ -87,7 +87,7 @@ func1 | SOFTPLUS '(' expr ')' # SoftplusFunc | GELU '(' expr ')' # GeluFunc | SIGN '(' expr ')' # SignFunc - + ; // Two-argument functions diff --git a/src/more_math/Parser/MathExprLexer.py b/src/more_math/Parser/MathExprLexer.py index edb4034..edd55ef 100644 --- a/src/more_math/Parser/MathExprLexer.py +++ b/src/more_math/Parser/MathExprLexer.py @@ -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 = [ "", - "'('", "')'", "','", "'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 = [ "", - "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" diff --git a/src/more_math/Parser/MathExprParser.py b/src/more_math/Parser/MathExprParser.py index b8445b8..7bd6846 100644 --- a/src/more_math/Parser/MathExprParser.py +++ b/src/more_math/Parser/MathExprParser.py @@ -155,28 +155,28 @@ class MathExprParser ( Parser ): sharedContextCache = PredictionContextCache() - literalNames = [ "", "'('", "')'", "','", "'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 = [ "", "'('", "')'", "','", "'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 = [ "", "", "", "", - "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 = [ "", "", "", "", + "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) - + diff --git a/src/more_math/helper_functions.py b/src/more_math/helper_functions.py index cd34c0d..15409ab 100644 --- a/src/more_math/helper_functions.py +++ b/src/more_math/helper_functions.py @@ -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) diff --git a/src/more_math/nodes.py b/src/more_math/nodes.py index 9998681..97a1037 100644 --- a/src/more_math/nodes.py +++ b/src/more_math/nodes.py @@ -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: diff --git a/tests/conftest.py b/tests/conftest.py index 12af46a..f481911 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 + diff --git a/tests/test_more_math.py b/tests/test_more_math.py index 86f5117..bea4c52 100644 --- a/tests/test_more_math.py +++ b/tests/test_more_math.py @@ -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...)