add bit count, bit shifts + tests + restore what AI deleted

This commit is contained in:
mcDandy
2026-02-13 23:15:50 +01:00
parent 93209bf261
commit 588a74cfeb
13 changed files with 3058 additions and 2292 deletions
+41 -2
View File
@@ -51,6 +51,7 @@ You can also get the node from comfy manager under the name of More math.
- Math: `+`, `-`, `*`, `/`, `%`, `^`, `|x|` (norm/abs)
- Boolean: `<`, `<=`, `>`, `>=`, `==`, `!=`
(`false = 0.0`, `true = 1.0`)
- Bitwise Shifts: `<<`, `>>` (left shift, right shift)
- Indexing: `x[i]` or `x[i, j, ...]` - Selects a sublist (if index count < number of dimensions) or value at position.
- Lists: `[v1, v2, ...]` (Vector math supported, mostly usefull in `conv` and `permute`)
- You can also use lists to do math with input tensor (image, noise, conditioing, latent, audio) which results in batched output as long as batch size is different to list size.
@@ -81,7 +82,7 @@ You can also get the node from comfy manager under the name of More math.
- `gamma(x)`: Gamma function.
- `dist(x1, y1, x2, y2)` or `distance`: Euclidean distance between points (x1, y1) and (x2, y2).
- `clamp(x, min, max)`: Constrains x to be between min and max.
- `step(x, edge)`: Returns 1.0 if x >= edge, else 0.0.
- `step(x, edge)`: Returns 1.0 if x ≥ edge, else 0.0.
### Trigonometric
@@ -131,7 +132,7 @@ You can also get the node from comfy manager under the name of More math.
- `moment(x, a, k)`: Returns the k-th moment of x centered around a.
- `topk(x, k)`: Returns a tensor with the **top K largest** values preserved at their original positions (others zeroed). For lists, returns the top K largest items sorted descending. (uses magnitude for complex numbers).
- `botk(x, k)`: Returns a tensor with the **bottom K smallest** values preserved at their original positions (others zeroed). For lists, returns the bottom K smallest items sorted ascending. (uses magnitude for complex numbers)
- `topk_ind(x, k)` or `topk_indices`: Returns the **indices** of the top K largest values in the flattened tensor.
- `topk_ind(x, k)` or `topk_indices: Returns the **indices** of the top K largest values in the flattened tensor.
- `botk_ind(x, k)` or `botk_indices`: Returns the **indices** of the bottom K smallest values in the flattened tensor.
- `sort(x)`: Sorts elements in ascending order along the last dimension.
- `argsort(x)` or `argsort(x, descending)`: Returns the **indices** that would sort the tensor/list. Optional second parameter for descending order.
@@ -218,6 +219,44 @@ Generates random noise with default shape of aither first input or maximum of in
- `random_bernoulli(seed, p,[shape])` or `randb`: generates a random tensor with Bernoulli distribution. Parameter `p` is the probability of getting 1, can be aither float or tensor. If p is tensor, shape is ignored.
- `random_poisson(seed, lambda,[shape])` or `randp`: generates a random tensor with Poisson distribution. Lambda can be either float or tensor.
### Bitwise Operations
Bitwise operations work with scalars, tensors, and lists, preserving bit patterns (especially important for floats where bit patterns are preserved, not values converted).
#### Shift Operators
- `a << b`: Left shift operator. Shifts bits of `a` left by `b` positions.
- `a >> b`: Right shift operator. Shifts bits of `a` right by `b` positions.
**Examples**:
```
5 << 2 # Returns: 20 (0b0101 << 2 = 0b10100)
20 >> 2 # Returns: 5 (0b10100 >> 2 = 0b0101)
[1, 2, 4] << 1 # Returns: [2, 4, 8]
```
#### Bitwise Functions
- `band(a, b)` or `bitwise_and(a, b)`: Bitwise AND. Returns bits set in both operands.
- `bor(a, b)` or `bitwise_or(a, b)`: Bitwise OR. Returns bits set in either operand.
- `xor(a, b)` or `bitwise_xor(a, b)`: Bitwise XOR. Returns bits set in exactly one operand.
- `bnot(a)` or `bitwise_not(a)`: Bitwise NOT. Inverts all bits in the operand.
- `bitcount(a)`, `popcount(a)`, or `popcnt(a)`: Count set bits. Returns the number of set bits (1s) in the binary representation as a float.
**Examples**:
```
band(5, 3) # Returns: 1 (0b0101 & 0b0011 = 0b0001)
bor(5, 3) # Returns: 7 (0b0101 | 0b0011 = 0b0111)
xor(5, 3) # Returns: 6 (0b0101 ^ 0b0011 = 0b0110)
bnot(5) # Returns: inverted bit pattern
bitcount(5) # Returns: 2.0 (0b0101 has 2 set bits)
bitcount(15) # Returns: 4.0 (0b1111 has 4 set bits)
```
**Bit Pattern Preservation**:
- **Integers**: Direct bitwise operations
- **Floats**: Bit patterns are preserved using `struct` module (pack float→int, operate, unpack int→float)
- **Tensors**: Uses `.view()` to reinterpret bytes without value conversion
- **Lists**: Element-wise operations applied to each element
### Stack
- `stack_push(id, value)`: Pushes value to stack with id.
+13 -4
View File
@@ -49,9 +49,14 @@ addExpr:
| mulExpr # ToMul;
mulExpr:
mulExpr MULT powExpr # MulExp
| mulExpr DIV powExpr # DivExp
| mulExpr MOD powExpr # ModExp
mulExpr MULT shiftExpr # MulExp
| mulExpr DIV shiftExpr # DivExp
| mulExpr MOD shiftExpr # ModExp
| shiftExpr # ToShift;
shiftExpr:
shiftExpr LSHIFT powExpr # LShiftExp
| shiftExpr RSHIFT powExpr # RShiftExp
| powExpr # ToPow;
powExpr: unaryExpr POW powExpr # PowExp | unaryExpr # ToUnary;
@@ -150,7 +155,8 @@ func1:
| FLATTEN LPAREN expr RPAREN # FlattenFunc
| MOTION_MASK LPAREN expr RPAREN # MotionMaskFunc
| FLOW_TO_IMAGE LPAREN expr RPAREN # FlowToImageFunc
| BNOT LPAREN expr RPAREN # BitNotFunc;
| BNOT LPAREN expr RPAREN # BitNotFunc
| BITCOUNT LPAREN expr RPAREN # BitCountFunc;
func2:
POWE LPAREN expr COMMA expr RPAREN # PowFunc
@@ -363,6 +369,7 @@ BAND: 'band';
BOR: 'bor';
XOR: 'xor';
BNOT: 'bnot';
BITCOUNT: 'bitcount' | 'popcount' | 'popcnt';
TENSOR: 'tensor';
@@ -372,6 +379,8 @@ MULT: '*';
DIV: '/';
MOD: '%';
POW: '^';
LSHIFT: '<<';
RSHIFT: '>>';
GE: '>=';
GT: '>';
File diff suppressed because one or more lines are too long
+63 -58
View File
@@ -127,38 +127,41 @@ BAND=126
BOR=127
XOR=128
BNOT=129
TENSOR=130
PLUS=131
MINUS=132
MULT=133
DIV=134
MOD=135
POW=136
GE=137
GT=138
LE=139
LT=140
EQ=141
EQUEALS=142
NE=143
PIPE=144
LPAREN=145
RPAREN=146
COMMA=147
SEMICOLON=148
ARROW=149
LBRACKET=150
RBRACKET=151
QUESTION=152
COLON=153
LBRACE=154
RBRACE=155
NUMBER=156
CONSTANT=157
VARIABLE=158
SL_COMMENT=159
ML_COMMENT=160
WS=161
BITCOUNT=130
TENSOR=131
PLUS=132
MINUS=133
MULT=134
DIV=135
MOD=136
POW=137
LSHIFT=138
RSHIFT=139
GE=140
GT=141
LE=142
LT=143
EQ=144
EQUEALS=145
NE=146
PIPE=147
LPAREN=148
RPAREN=149
COMMA=150
SEMICOLON=151
ARROW=152
LBRACKET=153
RBRACKET=154
QUESTION=155
COLON=156
LBRACE=157
RBRACE=158
NUMBER=159
CONSTANT=160
VARIABLE=161
SL_COMMENT=162
ML_COMMENT=163
WS=164
'sin'=1
'cos'=2
'tan'=3
@@ -256,29 +259,31 @@ WS=161
'bor'=127
'xor'=128
'bnot'=129
'tensor'=130
'+'=131
'-'=132
'*'=133
'/'=134
'%'=135
'^'=136
'>='=137
'>'=138
'<='=139
'<'=140
'=='=141
'='=142
'!='=143
'|'=144
'('=145
')'=146
','=147
';'=148
'->'=149
'['=150
']'=151
'?'=152
':'=153
'{'=154
'}'=155
'tensor'=131
'+'=132
'-'=133
'*'=134
'/'=135
'%'=136
'^'=137
'<<'=138
'>>'=139
'>='=140
'>'=141
'<='=142
'<'=143
'=='=144
'='=145
'!='=146
'|'=147
'('=148
')'=149
','=150
';'=151
'->'=152
'['=153
']'=154
'?'=155
':'=156
'{'=157
'}'=158
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+63 -58
View File
@@ -127,38 +127,41 @@ BAND=126
BOR=127
XOR=128
BNOT=129
TENSOR=130
PLUS=131
MINUS=132
MULT=133
DIV=134
MOD=135
POW=136
GE=137
GT=138
LE=139
LT=140
EQ=141
EQUEALS=142
NE=143
PIPE=144
LPAREN=145
RPAREN=146
COMMA=147
SEMICOLON=148
ARROW=149
LBRACKET=150
RBRACKET=151
QUESTION=152
COLON=153
LBRACE=154
RBRACE=155
NUMBER=156
CONSTANT=157
VARIABLE=158
SL_COMMENT=159
ML_COMMENT=160
WS=161
BITCOUNT=130
TENSOR=131
PLUS=132
MINUS=133
MULT=134
DIV=135
MOD=136
POW=137
LSHIFT=138
RSHIFT=139
GE=140
GT=141
LE=142
LT=143
EQ=144
EQUEALS=145
NE=146
PIPE=147
LPAREN=148
RPAREN=149
COMMA=150
SEMICOLON=151
ARROW=152
LBRACKET=153
RBRACKET=154
QUESTION=155
COLON=156
LBRACE=157
RBRACE=158
NUMBER=159
CONSTANT=160
VARIABLE=161
SL_COMMENT=162
ML_COMMENT=163
WS=164
'sin'=1
'cos'=2
'tan'=3
@@ -256,29 +259,31 @@ WS=161
'bor'=127
'xor'=128
'bnot'=129
'tensor'=130
'+'=131
'-'=132
'*'=133
'/'=134
'%'=135
'^'=136
'>='=137
'>'=138
'<='=139
'<'=140
'=='=141
'='=142
'!='=143
'|'=144
'('=145
')'=146
','=147
';'=148
'->'=149
'['=150
']'=151
'?'=152
':'=153
'{'=154
'}'=155
'tensor'=131
'+'=132
'-'=133
'*'=134
'/'=135
'%'=136
'^'=137
'<<'=138
'>>'=139
'>='=140
'>'=141
'<='=142
'<'=143
'=='=144
'='=145
'!='=146
'|'=147
'('=148
')'=149
','=150
';'=151
'->'=152
'['=153
']'=154
'?'=155
':'=156
'{'=157
'}'=158
+36
View File
@@ -296,6 +296,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#ToShift.
def enterToShift(self, ctx:MathExprParser.ToShiftContext):
pass
# Exit a parse tree produced by MathExprParser#ToShift.
def exitToShift(self, ctx:MathExprParser.ToShiftContext):
pass
# Enter a parse tree produced by MathExprParser#MulExp.
def enterMulExp(self, ctx:MathExprParser.MulExpContext):
pass
@@ -323,6 +332,24 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#RShiftExp.
def enterRShiftExp(self, ctx:MathExprParser.RShiftExpContext):
pass
# Exit a parse tree produced by MathExprParser#RShiftExp.
def exitRShiftExp(self, ctx:MathExprParser.RShiftExpContext):
pass
# Enter a parse tree produced by MathExprParser#LShiftExp.
def enterLShiftExp(self, ctx:MathExprParser.LShiftExpContext):
pass
# Exit a parse tree produced by MathExprParser#LShiftExp.
def exitLShiftExp(self, ctx:MathExprParser.LShiftExpContext):
pass
# Enter a parse tree produced by MathExprParser#ToPow.
def enterToPow(self, ctx:MathExprParser.ToPowContext):
pass
@@ -1115,6 +1142,15 @@ class MathExprListener(ParseTreeListener):
pass
# Enter a parse tree produced by MathExprParser#BitCountFunc.
def enterBitCountFunc(self, ctx:MathExprParser.BitCountFuncContext):
pass
# Exit a parse tree produced by MathExprParser#BitCountFunc.
def exitBitCountFunc(self, ctx:MathExprParser.BitCountFuncContext):
pass
# Enter a parse tree produced by MathExprParser#PowFunc.
def enterPowFunc(self, ctx:MathExprParser.PowFuncContext):
pass
File diff suppressed because it is too large Load Diff
+20
View File
@@ -169,6 +169,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToShift.
def visitToShift(self, ctx:MathExprParser.ToShiftContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#MulExp.
def visitMulExp(self, ctx:MathExprParser.MulExpContext):
return self.visitChildren(ctx)
@@ -184,6 +189,16 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#RShiftExp.
def visitRShiftExp(self, ctx:MathExprParser.RShiftExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#LShiftExp.
def visitLShiftExp(self, ctx:MathExprParser.LShiftExpContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#ToPow.
def visitToPow(self, ctx:MathExprParser.ToPowContext):
return self.visitChildren(ctx)
@@ -624,6 +639,11 @@ class MathExprVisitor(ParseTreeVisitor):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#BitCountFunc.
def visitBitCountFunc(self, ctx:MathExprParser.BitCountFuncContext):
return self.visitChildren(ctx)
# Visit a parse tree produced by MathExprParser#PowFunc.
def visitPowFunc(self, ctx:MathExprParser.PowFuncContext):
return self.visitChildren(ctx)
+84 -5
View File
@@ -113,15 +113,15 @@ class UnifiedMathVisitor(MathExprVisitor):
# one of them is a list and one is tensor
if self._is_tensor(a) and self._is_list(b):
if(a.shape[0]==len(b)):
A = torch.split(a,1)
if a.shape[0] == len(b):
A = torch.split(a, 1)
results = [self._bin_op(x.squeeze(0), y, torch_op, scalar_op) for x, y in zip(A, b)]
# Ensure all results are tensors
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
results = [self._bin_op(a, x, torch_op, scalar_op) for x in b]
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
results = [self._bin_op(a, x, torch_op, scalar_op) for x in b]
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
if self._is_list(a) and self._is_tensor(b):
if b.shape[0] == len(a):
B = torch.split(b, 1)
@@ -2172,6 +2172,19 @@ class UnifiedMathVisitor(MathExprVisitor):
error_msg = f"{ctx.start.line}:{ctx.start.column}: matmul({a.shape}, {b.shape}): {str(e)}"
raise ValueError(error_msg)
def visitToShift(self, ctx):
return (yield ctx.shiftExpr())
def visitLShiftExp(self, ctx):
a = yield ctx.shiftExpr()
b = yield ctx.powExpr()
return self._bitwise_op(a, b, torch.bitwise_left_shift, self._scalar_bitwise_lshift)
def visitRShiftExp(self, ctx):
a = yield ctx.shiftExpr()
b = yield ctx.powExpr()
return self._bitwise_op(a, b, torch.bitwise_right_shift, self._scalar_bitwise_rshift)
def visitBitAndFunc(self, ctx):
a = (yield ctx.expr(0))
b = (yield ctx.expr(1))
@@ -2191,6 +2204,10 @@ class UnifiedMathVisitor(MathExprVisitor):
v = (yield ctx.expr())
return self._bitwise_not(v)
def visitBitCountFunc(self, ctx):
v = (yield ctx.expr())
return self._bitwise_popcount(v)
def _bitwise_op(self, a, b, torch_op, scalar_op):
"""Binary bitwise operation handler supporting tensors, lists, and scalars."""
if self._is_tensor(a) and a.numel() == 1:
@@ -2309,3 +2326,65 @@ class UnifiedMathVisitor(MathExprVisitor):
return torch.int64
else:
return torch.int32
def _bitwise_popcount(self, v):
"""Count the number of set bits (1s) in the binary representation."""
if self._is_tensor(v):
v_t = self._promote_to_tensor(v).flatten().long()
# Use numpy's bin and count for efficiency
counts = torch.tensor([bin(int(x) & 0xFFFFFFFFFFFFFFFF).count('1') for x in v_t.tolist()],
dtype=torch.float32, device=v_t.device)
if counts.numel() == 1:
return float(counts.item())
return counts
if self._is_list(v):
return [self._bitwise_popcount(x) for x in v]
# Scalar - count set bits
v_int = int(v)
return float(bin(v_int & 0xFFFFFFFFFFFFFFFF).count('1'))
def _scalar_bitwise_lshift(self, a, b):
"""Scalar left shift with bit-pattern preservation for floats."""
b_int = int(b)
# If a is already an int, just do the shift
if isinstance(a, int):
return a << b_int
# For floats, preserve bit pattern
if isinstance(a, float):
fmt = 'd' # double (64-bit)
bit_fmt = 'Q' # unsigned long long
a_bits = struct.unpack(bit_fmt, struct.pack(fmt, a))[0]
result_bits = (a_bits << b_int) & ((1 << 64) - 1) # Mask to 64 bits
try:
return struct.unpack(fmt, struct.pack(bit_fmt, result_bits))[0]
except struct.error:
return float(result_bits & ((1 << 53) - 1)) # Return mantissa if error
# Fallback for other types
return int(a) << b_int
def _scalar_bitwise_rshift(self, a, b):
"""Scalar right shift with bit-pattern preservation for floats."""
b_int = int(b)
# If a is already an int, just do the shift
if isinstance(a, int):
return a >> b_int
# For floats, preserve bit pattern
if isinstance(a, float):
fmt = 'd' # double (64-bit)
bit_fmt = 'Q' # unsigned long long
a_bits = struct.unpack(bit_fmt, struct.pack(fmt, a))[0]
result_bits = a_bits >> b_int
try:
return struct.unpack(fmt, struct.pack(bit_fmt, result_bits))[0]
except struct.error:
return float(result_bits)
# Fallback for other types
return int(a) >> b_int
+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(f"✓ 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(f"✓ 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(f"✓ 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(f"✓ 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(f"✓ 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)