shape returns list now

This commit is contained in:
mcDandy
2026-02-17 00:34:51 +01:00
parent 0a8b87d574
commit 412ea5e89b
+6 -3
View File
@@ -2376,11 +2376,14 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitShapeFunc(self, ctx):
val = (yield ctx.expr())
if self._is_tensor(val):
return torch.tensor(list(val.shape), dtype=torch.long, device=self.device)
# Return shape as a 1D tensor of integers
return list(val.shape)
elif self._is_list(val):
return torch.tensor([len(val)], dtype=torch.long, device=self.device)
# Return list length as a single-element tensor
return [len(val)]
else:
return torch.tensor([], dtype=torch.long, device=self.device)
# Scalar has shape []
return []
def _bitwise_op(self, a, b, torch_op, scalar_op):
"""Binary bitwise operation handler supporting tensors, lists, and scalars."""