From 0d8cc9c2fe8f3c6d040f1049f5f2084e17351650 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Wed, 28 Jan 2026 19:09:15 +0100 Subject: [PATCH] fix longstanding bug of tensor OP list --- more_math/Parser/UnifiedMathVisitor.py | 8 ++++---- more_math/helper_functions.py | 4 ++-- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 4ae845d..8c1b927 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -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): diff --git a/more_math/helper_functions.py b/more_math/helper_functions.py index 3861251..39ae649 100644 --- a/more_math/helper_functions.py +++ b/more_math/helper_functions.py @@ -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):