add bit count, bit shifts + tests + restore what AI deleted
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
@@ -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
+656
-637
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+1761
-1526
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user