fix diagonal and better errors

This commit is contained in:
mcDandy
2026-08-26 14:41:14 +02:00
parent 4413d22e5e
commit 7d210b53f4
2 changed files with 21 additions and 9 deletions
+20 -5
View File
@@ -4959,7 +4959,8 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitDiagonalMatrixFunc(self, ctx):
"""diagonal_matrix(x) - create a diagonal matrix from a vector"""
x = self._promote_to_tensor((yield ctx.expr(0)))
shape = self._get_shape_from_ctx(ctx, 1)
shape_val = yield ctx.expr(1)
shape = self._normalize_shape_arg(shape_val, ctx, "diagonal_matrix shape")
offset = 0
dim1 = -2
dim2 = -1
@@ -4971,6 +4972,10 @@ class UnifiedMathVisitor(MathExprVisitor):
if len(ctx.expr()) > 4:
dim2 = yield ctx.expr(4)
offset = self._to_int(offset, ctx, "diagonal_matrix offset") if self._is_tensor(offset) or self._is_list(offset) else int(offset)
dim1 = self._to_int(dim1, ctx, "diagonal_matrix dim1") if self._is_tensor(dim1) or self._is_list(dim1) else int(dim1)
dim2 = self._to_int(dim2, ctx, "diagonal_matrix dim2") if self._is_tensor(dim2) or self._is_list(dim2) else int(dim2)
res = torch.zeros(shape, dtype=x.dtype, device=x.device)
diag_view = torch.diagonal(res, offset=offset, dim1=dim1, dim2=dim2)
diag_view.copy_(x)
@@ -4979,7 +4984,17 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitCoordsFunc(self, ctx):
"""coords(shape) - generate a grid of coordinates"""
shape = self._get_shape_from_ctx(ctx, 0)
dim = ctx.expr(1)
dtype = ctx.expr(2).dtype if len(ctx.expr()) > 2 else torch.float32
return getIndexTensorAlongDim(torch.zeros(shape, dtype=dtype, device=self.device), dim)
shape_val = yield ctx.expr(0)
shape = self._normalize_shape_arg(shape_val, ctx, "coords shape")
dim_val = yield ctx.expr(1)
dim = self._to_int(dim_val, ctx, "coords dim", strict=True)
dtype = torch.float32
if len(ctx.expr()) > 2:
dtype_val = yield ctx.expr(2)
if isinstance(dtype_val, torch.dtype):
dtype = dtype_val
base = torch.zeros(shape, dtype=dtype, device=self.device)
values = torch.arange(shape[dim], dtype=dtype, device=self.device)
view_shape = [1] * len(shape)
view_shape[dim] = shape[dim]
return values.view(*view_shape).expand(*shape)
+1 -4
View File
@@ -42,10 +42,7 @@ def parse_expr(expr: str):
def getIndexTensorAlongDim(tensor, dim):
"""Create a tensor of indices along a dimension, broadcasted to full shape."""
if tensor is torch.Tensor:
shape = tensor.shape
else:
shape = torch.Size(tensor)
shape = tensor.shape
values = torch.arange(shape[dim], dtype=torch.float32, device=tensor.device)
view_shape = [1] * len(shape)
view_shape[dim] = shape[dim]