diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index c9b3480..685f66f 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -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