From 1c32d60a5609c02732da4c84b09d7aaed53e684b Mon Sep 17 00:00:00 2001 From: mcDandy Date: Tue, 24 Feb 2026 20:13:19 +0100 Subject: [PATCH] add string support to [] --- more_math/Parser/UnifiedMathVisitor.py | 37 ++++++++++++++++++++++++-- 1 file changed, 35 insertions(+), 2 deletions(-) diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index 8369871..9dbcad8 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -248,6 +248,26 @@ class UnifiedMathVisitor(MathExprVisitor): # Use standard PyTorch/list indexing if self._is_tensor(val): + if len(indices) > val.ndim: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Expacted up to {val.ndim} dimensions but got {indices}.") + + for dim, idx in enumerate(indices): + if isinstance(idx, int): + size = val.shape[dim] + if idx < 0 or idx >= size: + raise ValueError( + f"{ctx.start.line}:{ctx.start.column}: Index {idx} out of bounds for dimension {dim} with size {size}" + ) + else: + idx_tensor = idx + if self._is_tensor(idx_tensor): + if idx_tensor.numel() == 0: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Empty tensor for dimension {dim}") + if torch.any(idx_tensor < 0) or torch.any(idx_tensor >= val.shape[dim]): + raise ValueError( + f"{ctx.start.line}:{ctx.start.column}: Index out of bounds for dimension {dim} with size {val.shape[dim]}" + ) + idx_tuple = tuple(indices) result = val[idx_tuple] if self._is_tensor(result): @@ -255,17 +275,30 @@ class UnifiedMathVisitor(MathExprVisitor): return result.item() return result.contiguous() return result + current = val + for idx in indices: + if isinstance(idx, torch.Tensor): + if idx.numel() != 1: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Too many indecies for string with 1 dimension. Got {idx.numel()}") + idx = int(idx.flatten()[0].item()) + if idx >= len(current) or idx < 0: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Index {idx} out of bounds for string of length {len(current)}") + current = current[idx] + return current elif self._is_list(val): # Navigate through nested lists current = val for idx in indices: if isinstance(idx, torch.Tensor): + if idx.numel() != 1: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: List index must be a scalar. Got tensor with length of {idx.numel()}") idx = int(idx.item()) + if idx >= len(current) or idx < 0: + raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Index {idx} out of bounds for list of length {len(current)}") current = current[idx] return current - error_prefix = f"{ctx.start.line}:{ctx.start.column}:" - raise ValueError(f"{error_prefix} Indexing only supported on tensors and lists (found {type(val).__name__})") + raise ValueError(f"{error_prefix} Indexing only supported on tensors, lists, and strings (found {type(val).__name__})") def visitToAtom(self, ctx): return (yield ctx.atom())