fix longstanding bug of tensor OP list

This commit is contained in:
mcDandy
2026-01-28 19:09:15 +01:00
parent 4ac582c045
commit 0d8cc9c2fe
2 changed files with 6 additions and 6 deletions
+4 -4
View File
@@ -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):
+2 -2
View File
@@ -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):