Add nested tensor creation + svd and way to reverse SVD

This commit is contained in:
mcDandy
2026-08-08 23:54:52 +02:00
parent c207f97e7c
commit e91b29defb
2 changed files with 56 additions and 1 deletions
+15 -1
View File
@@ -420,7 +420,14 @@ func1:
flow_angle(x) - computes the angle of optical flow (basically angle() but real and imaginery is separate in 2 dimensions)
*/
| FLOW_ANG LPAREN expr RPAREN # FlowAngFunc;
/**
as_nested(x) - converts a list to a nested tensor (special object containing tensors of different shapes behaving like a tensor)
*/
| AS_NESTED LPAREN expr RPAREN # AsNestedFunc;
/**
svd(x) - computes the singular value decomposition of x. Returns a list of 3 tensors: U, S, V such that x = matmul(matmul(U, diagonal_matrix(S,shape(U))), V)
*/
| SVD LPAREN expr RPAREN # SvdFunc;
func2:
/**
@@ -718,6 +725,10 @@ func5:
remap(v, i_min, i_max, o_min, o_max) - remaps values from input range to output range
*/
REMAP LPAREN expr COMMA expr COMMA expr COMMA expr COMMA expr RPAREN # RemapFunc;
/**
diagonal_matrix(x, shape, [offset], [dim1], [dim2]) - creates a diagonal matrix from x with specified shape. offset is the diagonal offset. dim1 and dim2 are the dimensions to place the diagonal on.
*/
| DIAG LPAREN expr COMMA expr (COMMA expr (COMMA expr (COMMA expr)?)?)? RPAREN # DiagonalMatrixFunc;
// N-argument functions
funcN:
@@ -973,6 +984,9 @@ INTERPOLATE_LINEAR: 'interpolate_linear';
INTERPOLATE_AREA: 'interpolate_area';
INTERPOLATE_NEAREST: 'interpolate_nearest' | 'interpolate_nearest_exact';
TEXT_IMAGE: 'text_image';
AS_NESTED: 'as_nested_tensor';
SVD: 'svd';
DIAG: 'diagonal_matrix';
IF: 'if';
ELSE: 'else';
+41
View File
@@ -1,5 +1,6 @@
import time
import os
from PIL.XVThumbImagePlugin import r
import torch
import math
import inspect
@@ -4673,3 +4674,43 @@ class UnifiedMathVisitor(MathExprVisitor):
)
target_shape = tuple(int(x) for x in shape)
return noise.view(target_shape)
def visitAsNestedFunc(self, ctx):
val = yield ctx.expr()
if self._is_list(val):
return NestedTensor(val)
if self._is_tensor(val):
return NestedTensor([val])
if isinstance(val, (int, float)):
return NestedTensor([torch.tensor(val, device=self.device)])
def visitSvdFunc(self, ctx):
"""svd(x) - singular value decomposition"""
x = self._promote_to_tensor((yield ctx.expr()))
if x.ndim < 2:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: svd expects at least 2D input, got shape {tuple(x.shape)}")
u, s, vh = torch.linalg.svd(x, full_matrices=True)
v = vh.mvector_transpose(-2, -1).conj()
result = [u, s, v]
return result
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)
offset = 0
dim1 = -2
dim2 = -1
if len(ctx.expr()) > 2:
offset = yield ctx.expr(2)
if len(ctx.expr()) > 3:
dim1 = yield ctx.expr(3)
if len(ctx.expr()) > 4:
dim2 = yield ctx.expr(4)
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)
return res