AI: rewrite bin_op and add fp16 support for bitwise operations
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user