fix longstanding bug of tensor OP list
This commit is contained in:
@@ -101,13 +101,13 @@ 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)):
|
||||
c = torch.split(a,1)
|
||||
return torch.cat([self._bin_op(x, y, torch_op, scalar_op) for x,y in zip(a,c)],dim=0)
|
||||
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)
|
||||
if self._is_list(a) and self._is_tensor(b):
|
||||
if(b.shape[0]==len(a)):
|
||||
c = torch.split(a,1)
|
||||
return torch.cat([self._bin_op(x, y, torch_op, scalar_op) for x,y in zip(c,b)],dim=0)
|
||||
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 self._is_list(a) and not self._is_tensor(b):
|
||||
|
||||
@@ -22,8 +22,8 @@ def as_tensor(value, shape):
|
||||
return value.contiguous()
|
||||
if isinstance(value, (float, int)):
|
||||
value = (value,)
|
||||
# If it's a scalar or list, broadcast to the reference shape provided.
|
||||
return torch.broadcast_to(torch.Tensor(value).to(dtype=torch.float32), shape).contiguous()
|
||||
return torch.broadcast_to(torch.Tensor(value).to(dtype=torch.float32), shape).contiguous()
|
||||
return torch.cat(value)
|
||||
|
||||
|
||||
def parse_expr(expr: str):
|
||||
|
||||
Reference in New Issue
Block a user