From bfe0fd213795e153ec01551d42ddecd004ff6060 Mon Sep 17 00:00:00 2001 From: mcDandy Date: Thu, 26 Feb 2026 23:19:40 +0100 Subject: [PATCH] fix wierd bug with reshape (AI) --- more_math/Parser/UnifiedMathVisitor.py | 24 ++++++++++++++++-------- 1 file changed, 16 insertions(+), 8 deletions(-) diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 1f15736..b565f31 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -176,6 +176,21 @@ class UnifiedMathVisitor(MathExprVisitor): return list_op(val) return val + def _to_int(self, x, ctx, context_name="operation"): + """Convert value to int, handling tensors and nested lists recursively""" + if self._is_tensor(x): + if x.numel() == 1: + return int(x.item()) + else: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {context_name} expects scalar dimensions, got tensor with shape {x.shape}") + elif self._is_list(x): + if len(x) == 1: + return self._to_int(x[0], ctx, context_name) + else: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {context_name} expects scalar dimensions, got list with {len(x)} elements") + else: + return int(float(x)) + # ======================== # Visitors # ======================== @@ -996,14 +1011,7 @@ class UnifiedMathVisitor(MathExprVisitor): elif isinstance(new_shape, (list, tuple)): result = [] for d in new_shape: - if self._is_tensor(d): - # Handle tensor elements in list - if d.numel() == 1: - result.append(int(d.item())) - else: - raise ValueError(f"{ctx.start.line}:{ctx.start.column}: reshape expects scalar dimensions, got tensor with shape {d.shape}") - else: - result.append(int(float(d))) + result.append(self._to_int(d, ctx, "reshape")) new_shape = result elif isinstance(new_shape, (int, float)): new_shape = [int(float(new_shape))]