AI refactor indexing

This commit is contained in:
mcDandy
2026-02-04 14:05:01 +01:00
parent f5e8c3ba8c
commit 8ffc9a394c
+65 -42
View File
@@ -208,23 +208,44 @@ class UnifiedMathVisitor(MathExprVisitor):
for e in ctx.expr():
index_args.append((yield e))
if self._is_tensor(val):
error_prefix = f"{ctx.start.line}:{ctx.start.column}:"
if self._is_tensor(val):
# Selection (tensor[list/tensor])
if len(index_args) == 1 and (self._is_list(index_args[0]) or self._is_tensor(index_args[0])):
idx = index_args[0]
if self._is_tensor(idx): idx = idx.flatten().tolist()
flat_indices = [int(i + val.shape[0] if i < 0 else i) for i in idx]
if self._is_tensor(idx):
idx = idx.flatten().tolist()
# Explicit bounds check for selection
max_idx = val.shape[0]
for i in idx:
if i < -max_idx or i >= max_idx:
raise IndexError(f"{error_prefix} Index {int(i)} out of bounds for tensor of size {max_idx}")
flat_indices = [int(i + max_idx if i < 0 else i) for i in idx]
idx_tensor = torch.tensor(flat_indices, device=self.device, dtype=torch.long)
return torch.index_select(val, 0, idx_tensor).contiguous()
# Tuple indexing
# Multi-dimensional indexing
torch_indices = []
for idx in index_args:
if self._is_list(idx):
torch_indices.append(torch.tensor(idx, device=self.device, dtype=torch.long))
elif self._is_tensor(idx):
torch_indices.append(idx.long())
if len(index_args) > val.ndim:
raise IndexError(f"{error_prefix} Too many indices for tensor of dimension {val.ndim}")
for i, idx in enumerate(index_args):
dim_size = val.shape[i]
if self._is_list(idx) or self._is_tensor(idx):
t_idx = torch.tensor(idx, device=self.device, dtype=torch.long) if self._is_list(idx) else idx.long()
# Bounds check for list/tensor indices in multi-dim
if torch.any(t_idx < -dim_size) or torch.any(t_idx >= dim_size):
# Find the first offending index for the error message
offender = t_idx[(t_idx < -dim_size) | (t_idx >= dim_size)][0].item()
raise IndexError(f"{error_prefix} Index {int(offender)} out of bounds for dimension {i} with size {dim_size}")
torch_indices.append(t_idx)
else:
# Scalar index
if idx < -dim_size or idx >= dim_size:
raise IndexError(f"{error_prefix} Index {int(idx)} out of bounds for dimension {i} with size {dim_size}")
torch_indices.append(int(idx))
res = val[tuple(torch_indices)]
@@ -235,21 +256,35 @@ class UnifiedMathVisitor(MathExprVisitor):
return float(res)
elif self._is_list(val):
# Nested list indexing or selection
if len(index_args) > 1:
curr = val
for idx in index_args:
curr = curr[int(idx + len(curr) if idx < 0 else idx)]
return curr
# Multi-level or selection indexing for lists
curr = val
for i, idx in enumerate(index_args):
if not self._is_list(curr):
raise TypeError(f"{error_prefix} Cannot index into non-list object at level {i}")
idx = index_args[0]
if self._is_list(idx) or self._is_tensor(idx):
if self._is_tensor(idx): idx = idx.flatten().tolist()
return [val[int(i + len(val) if i < 0 else i)] for i in idx]
else:
return val[int(idx + len(val) if idx < 0 else idx)]
list_len = len(curr)
if self._is_list(idx) or self._is_tensor(idx):
# List selection (only supported as the LAST index for now, or if it's the only one)
if i != len(index_args) - 1:
raise TypeError(f"{error_prefix} List selection only supported at the final indexing level")
if self._is_tensor(idx):
idx = idx.flatten().tolist()
for slice_idx in idx:
if slice_idx < -list_len or slice_idx >= list_len:
raise IndexError(f"{error_prefix} Index {int(slice_idx)} out of bounds for list of length {list_len}")
return [curr[int(j + list_len if j < 0 else j)] for j in idx]
else:
# Single index
if idx < -list_len or idx >= list_len:
raise IndexError(f"{error_prefix} Index {int(idx)} out of bounds for list of length {list_len}")
curr = curr[int(idx + list_len if idx < 0 else idx)]
return curr
else:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Indexing only supported on tensors and lists.")
raise ValueError(f"{error_prefix} Indexing only supported on tensors and lists (found {type(val).__name__})")
def visitToAtom(self, ctx):
return (yield ctx.atom())
@@ -1416,29 +1451,14 @@ class UnifiedMathVisitor(MathExprVisitor):
if b is None:
return a
if self._is_tensor(a) and a.numel() == 1:
if self._is_tensor(a) and a.ndim == 0:
a = a.item()
if self._is_tensor(b) and b.numel() == 1:
if self._is_tensor(b) and b.ndim == 0:
b = b.item()
if self._is_tensor(a) or self._is_tensor(b):
if self._is_tensor(a) and not self._is_tensor(b):
b_tensor = torch.tensor(b, device=self.device, dtype=a.dtype)
if a.ndim > 0:
target_shape = [1] + list(a.shape[1:])
b = b_tensor.expand(target_shape).contiguous()
else:
b = b_tensor
elif not self._is_tensor(a) and self._is_tensor(b):
a_tensor = torch.tensor(a, device=self.device, dtype=b.dtype)
if b.ndim > 0:
target_shape = [1] + list(b.shape[1:])
a = a_tensor.expand(target_shape).contiguous()
else:
a = a_tensor
else:
a = self._promote_to_tensor(a)
b = self._promote_to_tensor(b)
a = self._promote_to_tensor(a)
b = self._promote_to_tensor(b)
# Ensure at least 1D for cat
if a.ndim == 0:
@@ -1446,14 +1466,17 @@ class UnifiedMathVisitor(MathExprVisitor):
if b.ndim == 0:
b = b.unsqueeze(0)
if b.shape == a.shape[1:]:
b = b.unsqueeze(0)
return torch.cat((a, b), dim=0)
if not self._is_list(a):
a = [a]
if not self._is_list(b):
b = [b]
return a + b
return a + b
def visitStart(self, ctx):
from antlr4.tree.Tree import TerminalNode
count = ctx.getChildCount()