made linter happy by removing all trailing spaces and unused imports

This commit is contained in:
mcDandy
2025-12-11 18:33:27 +01:00
parent 4abd03eb94
commit 11b5ffcc3a
19 changed files with 484 additions and 431 deletions
+3 -45
View File
@@ -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
+4 -4
View File
@@ -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
View File
@@ -3,7 +3,7 @@
__all__ = [
"NODE_CLASS_MAPPINGS",
"NODE_id_MAPPINGS",
]
__author__ = """Daniel Martinek"""
-34
View File
@@ -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()
+2
View File
@@ -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
+2 -3
View File
@@ -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']}
+1 -1
View File
@@ -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):
"""
+3 -4
View File
@@ -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
+3 -4
View File
@@ -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 -2
View File
@@ -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
+1 -1
View File
@@ -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
+8 -8
View File
@@ -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.
+1 -1
View File
@@ -87,7 +87,7 @@ func1
| SOFTPLUS '(' expr ')' # SoftplusFunc
| GELU '(' expr ')' # GeluFunc
| SIGN '(' expr ')' # SignFunc
;
// Two-argument functions
+24 -25
View File
@@ -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"
+51 -51
View File
@@ -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)
+1 -1
View File
@@ -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 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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...)