AI: rewrite bin_op and add fp16 support for bitwise operations

This commit is contained in:
mcDandy
2026-02-13 16:23:03 +01:00
parent 2a0fa503c9
commit 0d65516cbb
+41 -24
View File
@@ -115,13 +115,22 @@ class UnifiedMathVisitor(MathExprVisitor):
if self._is_tensor(a) and self._is_list(b):
if(a.shape[0]==len(b)):
A = torch.split(a,1)
return torch.cat([self._bin_op(x, y, torch_op, scalar_op) for x,y in zip(A,b)],dim=0)
return torch.cat([self._bin_op(a, x, torch_op, scalar_op) for x in b], dim=0)
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)
if self._is_list(a) and self._is_tensor(b):
if(b.shape[0]==len(a)):
B = torch.split(b,1)
return torch.cat([self._bin_op(x, y, torch_op, scalar_op) for x,y in zip(a,B)],dim=0)
return torch.cat([self._bin_op(x, b, torch_op, scalar_op) for x in a], dim=0)
if b.shape[0] == len(a):
B = torch.split(b, 1)
results = [self._bin_op(x, y.squeeze(0), torch_op, scalar_op) for x, y in zip(a, 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(x, b, torch_op, scalar_op) for x in a]
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 not self._is_tensor(b):
if self._is_list(b):
@@ -133,6 +142,7 @@ class UnifiedMathVisitor(MathExprVisitor):
if not self._is_tensor(a) and self._is_list(b):
return [self._bin_op(a, x, torch_op, scalar_op) for x in b]
# Handle tensor operations
if self._is_tensor(a) or self._is_tensor(b):
if torch_op:
return torch_op(a, b).contiguous()
@@ -2192,13 +2202,21 @@ class UnifiedMathVisitor(MathExprVisitor):
if self._is_tensor(a) and self._is_list(b):
if a.shape[0] == len(b):
A = torch.split(a, 1)
return torch.cat([self._bitwise_op(x, y, torch_op, scalar_op) for x, y in zip(A, b)], dim=0)
return torch.cat([self._bitwise_op(a, x, torch_op, scalar_op) for x in b], dim=0)
results = [self._bitwise_op(x.squeeze(0), y, torch_op, scalar_op) for x, y in zip(A, 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._bitwise_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)
return torch.cat([self._bitwise_op(x, y, torch_op, scalar_op) for x, y in zip(a, B)], dim=0)
return torch.cat([self._bitwise_op(x, b, torch_op, scalar_op) for x in a], dim=0)
results = [self._bitwise_op(x, y.squeeze(0), torch_op, scalar_op) for x, y in zip(a, 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._bitwise_op(x, b, torch_op, scalar_op) for x in a]
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)
# Handle list-list and list-scalar combinations
if self._is_list(a) and not self._is_tensor(b):
@@ -2219,20 +2237,6 @@ class UnifiedMathVisitor(MathExprVisitor):
return scalar_op(a, b)
def _get_bitwise_view_dtype(self, elem_size):
"""Get appropriate integer dtype for bitwise operations based on element size."""
if elem_size == 1:
return torch.int8
elif elem_size == 2:
return torch.int16
elif elem_size == 4:
return torch.int32
elif elem_size == 8:
return torch.int64
else:
# Fallback for unusual sizes
return torch.int32
def _bitwise_not(self, v):
"""Unary bitwise NOT handling for tensors, lists and scalars with support for fp16 and int16."""
if self._is_tensor(v):
@@ -2263,3 +2267,16 @@ class UnifiedMathVisitor(MathExprVisitor):
return struct.unpack(fmt, struct.pack(bit_fmt, res_bits))[0]
except struct.error:
return int(res_bits)
def _get_bitwise_view_dtype(self, elem_size):
"""Get appropriate integer dtype for bitwise operations based on element size."""
if elem_size == 1:
return torch.int8
elif elem_size == 2:
return torch.int16
elif elem_size == 4:
return torch.int32
elif elem_size == 8:
return torch.int64
else:
return torch.int32