add compound assignments
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
+874
-855
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
+1627
-3375
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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");
|
||||
|
||||
Reference in New Issue
Block a user