add compound assignments

This commit is contained in:
mcDandy
2026-04-14 13:32:57 +02:00
parent 0e5ebde862
commit 17bb87ebf4
13 changed files with 2989 additions and 4310 deletions
+2 -1
View File
@@ -23,7 +23,7 @@ You can also get the node from comfy manager under the name of More math.
- Nodes for FLOAT, STRING, CONDITIONING, LATENT, IMAGE, MASK, NOISE, AUDIO, VIDEO, MODEL, CLIP, VAE, SIGMAS and GUIDER
- Vector Math: Support for List literals `[v1, v2, ...]` and operations between lists/scalars/tensors
- Custom functions `funcname(variable,variable,...)->expression;` they can be used in any later defined custom function or in expression. Shadowing inbuilt functions do not work. **Be careful with recursion. There is no stack limit. Got to 700 000 iterations before I got bored.**
- Custom variables `varname=expression;` They can be used in any later assigment or final expression.
- Custom variables `varname=expression;` They can be used in any later assigment or final expression. Compound assignments (`+=`, `-=`, `*=`, `/=`, `%=`) are also supported.
- Support for **indexed assignment**: `a[i, j, ...] = expression;`. Supports multidimensional tensors and nested lists.
- **Scalar Filling**: If the assigned value has only 1 element (scalar, 1-element list/tensor), it fills the entire selected slice.
- **Rank Matching**: Automatically squeezes leading ones from the value to match the rank of the target slice (e.g., assigning a 4D tensor with `dim0=1` to a 3D slice).
@@ -50,6 +50,7 @@ You can also get the node from comfy manager under the name of More math.
## Operators
- Math: `+`, `-`, `*`, `/`, `%`, `^`, `|x|` (norm/abs)
- Assignment: `=`, `+=`, `-=`, `*=`, `/=`, `%=`
- Boolean: `<`, `<=`, `>`, `>=`, `==`, `!=`
(`false = 0.0`, `true = 1.0`)
- Bitwise Shifts: `<<`, `>>` (left shift, right shift)
+6 -1
View File
@@ -7,7 +7,7 @@ funcDef:
VARIABLE LPAREN paramList? RPAREN ARROW (block | expr) SEMICOLON # FunctionDef;
varDef:
VARIABLE (LBRACKET expr (COMMA expr)* RBRACKET)* EQUEALS expr SEMICOLON;
VARIABLE (LBRACKET expr (COMMA expr)* RBRACKET)* (EQUEALS | PLUS_EQ | MINUS_EQ | MULT_EQ | DIV_EQ | MOD_EQ) expr SEMICOLON;
paramList: VARIABLE (COMMA VARIABLE)*;
@@ -464,6 +464,11 @@ LE: '<=';
LT: '<';
EQ: '==';
EQUEALS: '=';
PLUS_EQ: '+=';
MINUS_EQ: '-=';
MULT_EQ: '*=';
DIV_EQ: '/=';
MOD_EQ: '%=';
NE: '!=';
PIPE: '|';
LPAREN: '(';
File diff suppressed because one or more lines are too long
+43 -33
View File
@@ -182,26 +182,31 @@ LE=181
LT=182
EQ=183
EQUEALS=184
NE=185
PIPE=186
LPAREN=187
RPAREN=188
COMMA=189
SEMICOLON=190
ARROW=191
LBRACKET=192
RBRACKET=193
QUESTION=194
COLON=195
LBRACE=196
RBRACE=197
NUMBER=198
CONSTANT=199
STRING=200
VARIABLE=201
SL_COMMENT=202
ML_COMMENT=203
WS=204
PLUS_EQ=185
MINUS_EQ=186
MULT_EQ=187
DIV_EQ=188
MOD_EQ=189
NE=190
PIPE=191
LPAREN=192
RPAREN=193
COMMA=194
SEMICOLON=195
ARROW=196
LBRACKET=197
RBRACKET=198
QUESTION=199
COLON=200
LBRACE=201
RBRACE=202
NUMBER=203
CONSTANT=204
STRING=205
VARIABLE=206
SL_COMMENT=207
ML_COMMENT=208
WS=209
'sin'=1
'cos'=2
'tan'=3
@@ -342,16 +347,21 @@ WS=204
'<'=182
'=='=183
'='=184
'!='=185
'|'=186
'('=187
')'=188
','=189
';'=190
'->'=191
'['=192
']'=193
'?'=194
':'=195
'{'=196
'}'=197
'+='=185
'-='=186
'*='=187
'/='=188
'%='=189
'!='=190
'|'=191
'('=192
')'=193
','=194
';'=195
'->'=196
'['=197
']'=198
'?'=199
':'=200
'{'=201
'}'=202
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+43 -33
View File
@@ -182,26 +182,31 @@ LE=181
LT=182
EQ=183
EQUEALS=184
NE=185
PIPE=186
LPAREN=187
RPAREN=188
COMMA=189
SEMICOLON=190
ARROW=191
LBRACKET=192
RBRACKET=193
QUESTION=194
COLON=195
LBRACE=196
RBRACE=197
NUMBER=198
CONSTANT=199
STRING=200
VARIABLE=201
SL_COMMENT=202
ML_COMMENT=203
WS=204
PLUS_EQ=185
MINUS_EQ=186
MULT_EQ=187
DIV_EQ=188
MOD_EQ=189
NE=190
PIPE=191
LPAREN=192
RPAREN=193
COMMA=194
SEMICOLON=195
ARROW=196
LBRACKET=197
RBRACKET=198
QUESTION=199
COLON=200
LBRACE=201
RBRACE=202
NUMBER=203
CONSTANT=204
STRING=205
VARIABLE=206
SL_COMMENT=207
ML_COMMENT=208
WS=209
'sin'=1
'cos'=2
'tan'=3
@@ -342,16 +347,21 @@ WS=204
'<'=182
'=='=183
'='=184
'!='=185
'|'=186
'('=187
')'=188
','=189
';'=190
'->'=191
'['=192
']'=193
'?'=194
':'=195
'{'=196
'}'=197
'+='=185
'-='=186
'*='=187
'/='=188
'%='=189
'!='=190
'|'=191
'('=192
')'=193
','=194
';'=195
'->'=196
'['=197
']'=198
'?'=199
':'=200
'{'=201
'}'=202
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -1,4 +1,4 @@
# Generated from ./MathExpr.g4 by ANTLR 4.13.2
# Generated from MathExpr.g4 by ANTLR 4.13.2
from antlr4 import *
if "." in __name__:
from .MathExprParser import MathExprParser
+58 -4
View File
@@ -1791,9 +1791,32 @@ class UnifiedMathVisitor(MathExprVisitor):
var_name = ctx.VARIABLE().getText()
expr_list = ctx.expr()
assign_op = "="
if getattr(ctx, "PLUS_EQ", lambda: None)() is not None: assign_op = "+="
elif getattr(ctx, "MINUS_EQ", lambda: None)() is not None: assign_op = "-="
elif getattr(ctx, "MULT_EQ", lambda: None)() is not None: assign_op = "*="
elif getattr(ctx, "DIV_EQ", lambda: None)() is not None: assign_op = "/="
elif getattr(ctx, "MOD_EQ", lambda: None)() is not None: assign_op = "%="
if not ctx.LBRACKET():
# Standard assignment: x = value
val = yield expr_list[0]
if assign_op != "=":
if var_name not in self.variables:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Variable '{var_name}' not defined for compound assignment.")
existing_val = self.variables[var_name]
if assign_op == "+=":
val = self._bin_op(existing_val, val, torch.add, lambda a, b: a + b, ctx)
elif assign_op == "-=":
val = self._bin_op(existing_val, val, torch.sub, lambda a, b: a - b, ctx)
elif assign_op == "*=":
val = self._bin_op(existing_val, val, torch.mul, lambda a, b: a * b, ctx)
elif assign_op == "/=":
val = self._bin_op(existing_val, val, torch.div, lambda a, b: a / b, ctx)
elif assign_op == "%=":
val = self._bin_op(existing_val, val, torch.remainder, lambda a, b: a % b, ctx)
self.variables[var_name] = val
return val
@@ -1830,13 +1853,27 @@ class UnifiedMathVisitor(MathExprVisitor):
# Target slice - used to compute expected shape
target_slice = target[idx_tuple]
if assign_op != "=":
if assign_op == "+=":
val_t = self._bin_op(target_slice, val_t, torch.add, lambda a, b: a + b, ctx)
elif assign_op == "-=":
val_t = self._bin_op(target_slice, val_t, torch.sub, lambda a, b: a - b, ctx)
elif assign_op == "*=":
val_t = self._bin_op(target_slice, val_t, torch.mul, lambda a, b: a * b, ctx)
elif assign_op == "/=":
val_t = self._bin_op(target_slice, val_t, torch.div, lambda a, b: a / b, ctx)
elif assign_op == "%=":
val_t = self._bin_op(target_slice, val_t, torch.remainder, lambda a, b: a % b, ctx)
val_t = self._promote_to_tensor(val_t) # Ensure it's still a tensor
# Squeeze leading ones to match target slice rank if it's smaller
# but target_slice.ndim might be 0 if it's a scalar location.
while val_t.ndim > target_slice.ndim and val_t.shape[0] == 1:
val_t = val_t.squeeze(0)
target[idx_tuple] = val_t
return assigned_val
return assigned_val if assign_op == "=" else val_t
except Exception as e:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Indexed assignment to '{var_name}' failed: {str(e)}")
elif self._is_list(target):
@@ -1845,8 +1882,25 @@ class UnifiedMathVisitor(MathExprVisitor):
for idx in indices[:-1]:
curr = curr[int(idx + len(curr) if idx < 0 else idx)]
last_idx = int(indices[-1])
curr[last_idx + len(curr) if last_idx < 0 else last_idx] = assigned_val
return assigned_val
real_idx = last_idx + len(curr) if last_idx < 0 else last_idx
if assign_op != "=":
existing_val = curr[real_idx]
if assign_op == "+=":
new_val = self._bin_op(existing_val, assigned_val, torch.add, lambda a, b: a + b, ctx)
elif assign_op == "-=":
new_val = self._bin_op(existing_val, assigned_val, torch.sub, lambda a, b: a - b, ctx)
elif assign_op == "*=":
new_val = self._bin_op(existing_val, assigned_val, torch.mul, lambda a, b: a * b, ctx)
elif assign_op == "/=":
new_val = self._bin_op(existing_val, assigned_val, torch.div, lambda a, b: a / b, ctx)
elif assign_op == "%=":
new_val = self._bin_op(existing_val, assigned_val, torch.remainder, lambda a, b: a % b, ctx)
curr[real_idx] = new_val
return new_val
else:
curr[real_idx] = assigned_val
return assigned_val
else:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Indexed assignment not supported for {type(target)}")
@@ -3405,7 +3459,7 @@ class UnifiedMathVisitor(MathExprVisitor):
return res
tensors = [self._promote_to_tensor(x) for x in items]
d = int(dim_val.item()) if self._is_tensor(dim_val) else int(dim_val)
d = int(dim_val.item()) if self._is_tensor(dim_val) else int(dim_val)
return torch.cat(tensors, dim=d)
def visitIntFunc(self, ctx):
+133
View File
@@ -0,0 +1,133 @@
#!/usr/bin/env python3
"""Quick test to verify bitwise operations work with tensors"""
import sys
sys.path.insert(0, r'D:\stability\Data\Packages\ComfyUI')
import torch
from custom_nodes.more_math.more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor
def test_bitwise_xor_float():
"""Test XOR with float tensors (the reported bug)"""
print("Testing bitwise XOR with float tensors...")
# Create float tensors (this was causing the error)
a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32)
b = torch.tensor([4.0, 5.0, 6.0], dtype=torch.float32)
visitor = UnifiedMathVisitor({"a": a, "b": b})
try:
# This should convert to int64, perform XOR, then convert back
result = visitor._bitwise_op(a, b, torch.bitwise_xor, lambda x, y: x ^ y)
print("✓ XOR succeeded!")
print(f" Input a (float32): {a}")
print(f" Input b (float32): {b}")
print(f" Result: {result}")
print(f" Result dtype: {result.dtype}")
return True
except Exception as e:
print(f"✗ XOR failed: {e}")
return False
def test_bitwise_and_float():
"""Test AND with float tensors"""
print("\nTesting bitwise AND with float tensors...")
a = torch.tensor([15.0, 14.0, 13.0], dtype=torch.float32)
b = torch.tensor([7.0, 3.0, 1.0], dtype=torch.float32)
visitor = UnifiedMathVisitor({"a": a, "b": b})
try:
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
print("✓ AND succeeded!")
print(f" Input a (float32): {a}")
print(f" Input b (float32): {b}")
print(f" Result: {result}")
print(f" Result dtype: {result.dtype}")
return True
except Exception as e:
print(f"✗ AND failed: {e}")
return False
def test_bitwise_or_float():
"""Test OR with float tensors"""
print("\nTesting bitwise OR with float tensors...")
a = torch.tensor([15.0, 14.0, 13.0], dtype=torch.float32)
b = torch.tensor([7.0, 3.0, 1.0], dtype=torch.float32)
visitor = UnifiedMathVisitor({"a": a, "b": b})
try:
result = visitor._bitwise_op(a, b, torch.bitwise_or, lambda x, y: x | y)
print("✓ OR succeeded!")
print(f" Input a (float32): {a}")
print(f" Input b (float32): {b}")
print(f" Result: {result}")
print(f" Result dtype: {result.dtype}")
return True
except Exception as e:
print(f"✗ OR failed: {e}")
return False
def test_bitwise_not_float():
"""Test NOT with float tensors"""
print("\nTesting bitwise NOT with float tensors...")
a = torch.tensor([1.0, 2.0, 3.0], dtype=torch.float32)
visitor = UnifiedMathVisitor({"a": a})
try:
result = visitor._bitwise_not(a)
print("✓ NOT succeeded!")
print(f" Input a (float32): {a}")
print(f" Result: {result}")
print(f" Result dtype: {result.dtype}")
return True
except Exception as e:
print(f"✗ NOT failed: {e}")
return False
def test_int_tensors_preserved():
"""Ensure int tensors still work as before"""
print("\nTesting that int tensor dtypes are preserved...")
a = torch.tensor([15, 14, 13], dtype=torch.int16)
b = torch.tensor([7, 3, 1], dtype=torch.int16)
visitor = UnifiedMathVisitor({"a": a, "b": b})
try:
result = visitor._bitwise_op(a, b, torch.bitwise_and, lambda x, y: x & y)
assert result.dtype == torch.int16, f"Expected int16, got {result.dtype}"
print("✓ Int16 dtype preserved!")
print(f" Input a (int16): {a}")
print(f" Input b (int16): {b}")
print(f" Result (int16): {result}")
return True
except Exception as e:
print(f"✗ Int16 test failed: {e}")
return False
if __name__ == "__main__":
print("=" * 70)
print("Bitwise Operations Fix Verification")
print("=" * 70)
results = [
test_bitwise_xor_float(),
test_bitwise_and_float(),
test_bitwise_or_float(),
test_bitwise_not_float(),
test_int_tensors_preserved(),
]
print("\n" + "=" * 70)
if all(results):
print("✓ All tests passed!")
else:
print("✗ Some tests failed")
sys.exit(1)
print("=" * 70)
+170
View File
@@ -0,0 +1,170 @@
#!/usr/bin/env python3
"""
Test bitwise shift operators and bit count function
"""
import torch
import sys
sys.path.insert(0, 'custom_nodes/more_math')
from more_math.Parser.MathExprParser import MathExprParser
from more_math.Parser.UnifiedMathVisitor import UnifiedMathVisitor
from antlr4 import InputStream, CommonTokenFactory
def parse_and_evaluate(expression, variables=None):
"""Parse and evaluate a math expression"""
if variables is None:
variables = {}
input_stream = InputStream(expression)
lexer = MathExprParser(input_stream).lexer
stream = CommonTokenFactory()
parser = MathExprParser(input_stream)
tree = parser.start()
visitor = UnifiedMathVisitor(variables, device='cpu')
result = visitor.visit(tree)
return result
def test_bit_shifts():
"""Test bitwise shift operators"""
print("=" * 60)
print("Testing Bitwise Shift Operators")
print("=" * 60)
test_cases = [
# Left shift: 5 << 2 = 20 (0101 << 2 = 10100)
("5 << 2", {}, 20),
# Right shift: 20 >> 2 = 5 (10100 >> 2 = 0101)
("20 >> 2", {}, 5),
# Left shift with variable
("x << 3", {"x": 4}, 32), # 4 << 3 = 32
# Right shift with variable
("x >> 2", {"x": 16}, 4), # 16 >> 2 = 4
# Chained shifts
("(8 << 2) >> 3", {}, 4), # (32) >> 3 = 4
]
for expr, vars, expected in test_cases:
try:
result = parse_and_evaluate(expr, vars)
status = "✓" if result == expected else "✗"
print(f"{status} {expr:30} = {result:10} (expected {expected})")
except Exception as e:
print(f"✗ {expr:30} ERROR: {e}")
def test_bit_count():
"""Test bit count function"""
print("\n" + "=" * 60)
print("Testing Bitwise Bit Count Function")
print("=" * 60)
test_cases = [
# bitcount(5) = 2 (0101 has 2 set bits)
("bitcount(5)", {}, 2),
# bitcount(15) = 4 (1111 has 4 set bits)
("bitcount(15)", {}, 4),
# bitcount(7) = 3 (111 has 3 set bits)
("bitcount(7)", {}, 3),
# bitcount(255) = 8 (11111111 has 8 set bits)
("bitcount(255)", {}, 8),
# bitcount(0) = 0
("bitcount(0)", {}, 0),
# With variable
("bitcount(x)", {"x": 31}, 5), # 31 = 11111 = 5 bits set
]
for expr, vars, expected in test_cases:
try:
result = parse_and_evaluate(expr, vars)
status = "✓" if result == expected else "✗"
print(f"{status} {expr:30} = {result:10} (expected {expected})")
except Exception as e:
print(f"✗ {expr:30} ERROR: {e}")
def test_bit_shifts_with_tensors():
"""Test bitwise shifts with tensors"""
print("\n" + "=" * 60)
print("Testing Bitwise Shifts with Tensors")
print("=" * 60)
vars = {
"a": torch.tensor([1, 2, 4, 8], dtype=torch.int32),
"shift": 2,
}
try:
result = parse_and_evaluate("a << shift", vars)
expected = torch.tensor([4, 8, 16, 32], dtype=torch.int32)
match = torch.equal(result, expected)
status = "✓" if match else "✗"
print(f"{status} tensor_shift_left: [1,2,4,8] << 2 = {result.tolist()}")
except Exception as e:
print(f"✗ tensor_shift_left ERROR: {e}")
try:
result = parse_and_evaluate("a >> shift", vars)
expected = torch.tensor([0, 0, 1, 2], dtype=torch.int32)
match = torch.equal(result, expected)
status = "✓" if match else "✗"
print(f"{status} tensor_shift_right: [1,2,4,8] >> 2 = {result.tolist()}")
except Exception as e:
print(f"✗ tensor_shift_right ERROR: {e}")
def test_bit_count_with_tensors():
"""Test bit count with tensors"""
print("\n" + "=" * 60)
print("Testing Bit Count with Tensors")
print("=" * 60)
vars = {
"nums": torch.tensor([5, 15, 7, 255], dtype=torch.int32),
}
try:
result = parse_and_evaluate("bitcount(nums)", vars)
# Expected: [2, 4, 3, 8] set bits
print(f"✓ bitcount([5,15,7,255]): {result}")
except Exception as e:
print(f"✗ bitcount_tensor ERROR: {e}")
def test_combinations():
"""Test combinations of shift and bit count"""
print("\n" + "=" * 60)
print("Testing Combinations")
print("=" * 60)
test_cases = [
# Shift then count bits
("bitcount(5 << 2)", {}, 2), # 5 << 2 = 20 (10100) = 2 bits
# Combined with other operators
("(5 << 2) | (3 << 4)", {}, 0xCC), # 20 | 48 = 0xCC = 204
]
for expr, vars, expected in test_cases:
try:
result = parse_and_evaluate(expr, vars)
status = "✓" if result == expected else "✗"
print(f"{status} {expr:35} = {result:10} (expected {expected})")
except Exception as e:
print(f"✗ {expr:35} ERROR: {e}")
if __name__ == "__main__":
test_bit_shifts()
test_bit_count()
test_bit_shifts_with_tensors()
test_bit_count_with_tensors()
test_combinations()
print("\n" + "=" * 60)
print("All tests completed!")
print("=" * 60)
+5 -5
View File
@@ -20,7 +20,7 @@ const FUNCTIONS = new Set([
"bnot", "bitwise_not", "bitcount", "popcount", "popcnt", "shape", "band", "bitwise_and", "bxor", "bitwise_xor",
"bor", "bitwise_or", "tensor", "stack_push", "stack_pop", "stack_clear", "stack_has", "stack_get", "timestamp","now",
"sort", "argsort", "argmin", "argmax", "softmax", "softmin", "unique", "flip", "cov", "corr", "correlation", "entropy",
"crop", "cat", "concatenate", "concat", "float","int", "linspace", "logspace", "roll",
"crop", "cat", "concatenate", "concat", "float","int", "linspace", "logspace", "roll","select",
"noise", "randn", "random_normal", "rand", "randu", "random_uniform", "randc", "random_cauchy", "rande",
"random_exponential", "randln", "random_log_normal", "randb", "random_bernoulli", "randp", "random_poisson", "randg",
"random_gamma", "randbeta", "random_beta", "randl", "random_laplace", "randgumbel", "random_gumbel", "randw",
@@ -60,7 +60,7 @@ function escapeHtml(value) {
function tokenize(text) {
const tokens = [];
const pattern = /#.*|\/\*[\s\S]*?\*\/|"(?:\\.|[^"\\\r\n])*"|'(?:\\.|[^'\\\r\n])*'|\b\d+(?:\.\d*)?(?:[eE][+-]?\d+)?\b|\B\.\d+(?:[eE][+-]?\d+)?\b|==|!=|>=|<=|<<|>>|->|[+\-*/%^=<>|?:,;()\[\]{}]|\b[a-zA-Z_][a-zA-Z_0-9]*\b|\s+|./g;
const pattern = /#.*|\/\*[\s\S]*?\*\/|"(?:\\.|[^"\\\r\n])*"|'(?:\\.|[^'\\\r\n])*'|\b\d+(?:\.\d*)?(?:[eE][+-]?\d+)?\b|\B\.\d+(?:[eE][+-]?\d+)?\b|==|!=|>=|<=|<<|>>|->|\+=|-=|\*=|\/=|%=|[+\-*/%^=<>|?:,;()\[\]{}]|\b[a-zA-Z_][a-zA-Z_0-9]*\b|\s+|./g;
let match;
const bracketStack = [];
const depthByType = { '(': 0, '[': 0, '{': 0 };
@@ -170,8 +170,8 @@ function ensureStyles() {
display: flex;
align-items: stretch;
width: 100%;
height: 100%;
min-height: 100%;
height: 100%;
min-height: 100%;
gap: 0;
overflow: hidden;
}
@@ -297,7 +297,7 @@ function attachLineNumbers(widget) {
editorContainer.appendChild(syntaxLayer);
editorContainer.appendChild(inputEl);
wrapper.style.height = "100%";
wrapper.style.height = "100%";
parent.style.height = "100%";
inputEl.classList.add("mrmth-line-input");