add ctx to bin_op to fix errors

This commit is contained in:
mcDandy
2026-02-26 21:18:50 +01:00
parent 6d0d804a9e
commit a9ed8cae2f
+71 -71
View File
@@ -552,16 +552,16 @@ class UnifiedMathVisitor(MathExprVisitor):
# Two-argument functions
def visitPowFunc(self, ctx):
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.pow, math.pow)
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.pow, math.pow,ctx)
def visitAtan2Func(self, ctx):
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.atan2, math.atan2)
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.atan2, math.atan2,ctx)
def visitTMinFunc(self, ctx):
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.minimum, min)
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.minimum, min,ctx)
def visitTMaxFunc(self, ctx):
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.maximum, max)
return self._bin_op((yield ctx.expr(0)), (yield ctx.expr(1)), torch.maximum, max,ctx)
def visitStepFunc(self, ctx):
# step(x, edge) = 1 if x >= edge else 0
@@ -569,7 +569,7 @@ class UnifiedMathVisitor(MathExprVisitor):
(yield ctx.expr(0)),
(yield ctx.expr(1)),
lambda x, edge: torch.where(x >= edge, 1.0, 0.0),
lambda x, edge: 1.0 if x >= edge else 0.0,
lambda x, edge: 1.0 if x >= edge else 0.0, ctx
)
def visitTopkFunc(self, ctx):
@@ -1279,7 +1279,7 @@ class UnifiedMathVisitor(MathExprVisitor):
a = (yield ctx.expr(1))
k = (yield ctx.expr(2))
return float(torch.sum(self._bin_op(self._bin_op(x,a,torch.sub,lambda x, a: x - a),k,torch.pow,pow)).item())/x.numel()
return float(torch.sum(self._bin_op(self._bin_op(x,a,torch.sub,lambda x, a: x - a,ctx),k,torch.pow,pow,ctx)).item())/x.numel()
def visitSortFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
@@ -2180,10 +2180,10 @@ class UnifiedMathVisitor(MathExprVisitor):
tsr = self._promote_to_tensor(tsr_val)
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
kernel_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
original_shape = tsr.shape
tsr = tsr.float()
reshap = False
if len(ctx.expr()) >= 2:
reshap_val = yield ctx.expr(1)
@@ -2974,21 +2974,21 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitSplitFunc(self, ctx):
string = yield ctx.expr(0)
delimiter = yield ctx.expr(1) if len(ctx.expr()) > 1 else " "
if not isinstance(string, str):
string = str(string)
if not isinstance(delimiter, str):
delimiter = str(delimiter)
return string.split(delimiter)
def visitJoinFunc(self, ctx):
items = yield ctx.expr(0)
separator = yield ctx.expr(1) if len(ctx.expr()) > 1 else ""
if not isinstance(separator, str):
separator = str(separator)
if self._is_list(items):
return separator.join([str(x) for x in items])
elif self._is_tensor(items):
@@ -3000,12 +3000,12 @@ class UnifiedMathVisitor(MathExprVisitor):
string = yield ctx.expr(0)
start = yield ctx.expr(1)
length = yield ctx.expr(2) if len(ctx.expr()) > 2 else None
if not isinstance(string, str):
string = str(string)
start_idx = int(start.item()) if self._is_tensor(start) else int(start)
if length is not None:
length_val = int(length.item()) if self._is_tensor(length) else int(length)
return string[start_idx:start_idx + length_val]
@@ -3015,12 +3015,12 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitFindFunc(self, ctx):
string = yield ctx.expr(0)
search = yield ctx.expr(1)
if not isinstance(string, str):
string = str(string)
if not isinstance(search, str):
search = str(search)
return float(string.find(search))
def visitTrimFunc(self, ctx):
@@ -3034,21 +3034,21 @@ class UnifiedMathVisitor(MathExprVisitor):
tsr = self._promote_to_tensor(tsr_val)
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
kernel_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
original_shape = tsr.shape
tsr = tsr.float()
def dilate_op(x):
kernel = torch.ones((kernel_size, kernel_size), device=x.device, dtype=x.dtype)
kernel = kernel.unsqueeze(0).unsqueeze(0)
kernel = kernel.repeat(x.size(1), 1, 1, 1)
pad = kernel_size // 2
x_padded = F.pad(x, (pad, pad, pad, pad), mode='replicate')
result = F.conv2d(x_padded, kernel, padding=0, groups=x.size(1))
return torch.clamp(result, 0, 1)
return self._apply_spatial_op(tsr, dilate_op, original_shape)
def visitErodeFunc(self, ctx):
@@ -3056,88 +3056,88 @@ class UnifiedMathVisitor(MathExprVisitor):
tsr = self._promote_to_tensor(tsr_val)
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
kernel_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
original_shape = tsr.shape
tsr = tsr.float()
def erode_op(x):
x_inv = 1.0 - x
kernel = torch.ones((kernel_size, kernel_size), device=x.device, dtype=x.dtype)
kernel = kernel.unsqueeze(0).unsqueeze(0)
kernel = kernel.repeat(x.size(1), 1, 1, 1)
pad = kernel_size // 2
x_padded = F.pad(x_inv, (pad, pad, pad, pad), mode='replicate')
result = F.conv2d(x_padded, kernel, padding=0, groups=x.size(1))
return torch.clamp(1.0 - result, 0, 1)
return self._apply_spatial_op(tsr, erode_op, original_shape)
def visitMorphOpenFunc(self, ctx):
tsr_val = yield ctx.expr(0)
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
eroded = yield from self.visitErodeFunc(ctx)
tsr = self._promote_to_tensor(eroded)
k_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
original_shape = tsr.shape
tsr = tsr.float()
def dilate_op(x):
kernel = torch.ones((k_size, k_size), device=x.device, dtype=x.dtype)
kernel = kernel.unsqueeze(0).unsqueeze(0)
kernel = kernel.repeat(x.size(1), 1, 1, 1)
pad = k_size // 2
x_padded = F.pad(x, (pad, pad, pad, pad), mode='replicate')
result = F.conv2d(x_padded, kernel, padding=0, groups=x.size(1))
return torch.clamp(result, 0, 1)
return self._apply_spatial_op(tsr, dilate_op, original_shape)
def visitMorphCloseFunc(self, ctx):
tsr_val = yield ctx.expr(0)
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
dilated = yield from self.visitDilateFunc(ctx)
tsr = self._promote_to_tensor(dilated)
k_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
original_shape = tsr.shape
tsr = tsr.float()
def erode_op(x):
x_inv = 1.0 - x
kernel = torch.ones((k_size, k_size), device=x.device, dtype=x.dtype)
kernel = kernel.unsqueeze(0).unsqueeze(0)
kernel = kernel.repeat(x.size(1), 1, 1, 1)
pad = k_size // 2
x_padded = F.pad(x_inv, (pad, pad, pad, pad), mode='replicate')
result = F.conv2d(x_padded, kernel, padding=0, groups=x.size(1))
return torch.clamp(1.0 - result, 0, 1)
return self._apply_spatial_op(tsr, erode_op, original_shape)
def visitRgbToHsvFunc(self, ctx):
num_args = len(ctx.expr())
# Determine mode: 1=tensor, 2=tensor+degrees, 3=r,g,b, 4=r,g,b+degrees
if num_args == 1 or num_args == 2:
# Tensor mode
rgb = self._promote_to_tensor((yield ctx.expr(0)))
if rgb.shape[-1] != 3:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: rgb_to_hsv expects tensor with last dim=3, got shape {rgb.shape}")
r = rgb[..., 0]
g = rgb[..., 1]
b = rgb[..., 2]
use_degrees = False
if num_args == 2:
degrees_val = yield ctx.expr(1)
@@ -3147,44 +3147,44 @@ class UnifiedMathVisitor(MathExprVisitor):
r = self._promote_to_tensor((yield ctx.expr(0)))
g = self._promote_to_tensor((yield ctx.expr(1)))
b = self._promote_to_tensor((yield ctx.expr(2)))
use_degrees = False
if num_args == 4:
degrees_val = yield ctx.expr(3)
use_degrees = bool(degrees_val.item() if self._is_tensor(degrees_val) else degrees_val)
# Clamp to [0, 1]
r = torch.clamp(r, 0, 1)
g = torch.clamp(g, 0, 1)
b = torch.clamp(b, 0, 1)
# RGB to HSV conversion
max_rgb, _ = torch.max(torch.stack([r, g, b]), dim=0)
min_rgb, _ = torch.min(torch.stack([r, g, b]), dim=0)
diff = max_rgb - min_rgb
# Hue (in degrees 0-360)
h = torch.zeros_like(max_rgb)
mask_r = (max_rgb == r) & (diff > 0)
h[mask_r] = (60 * ((g[mask_r] - b[mask_r]) / diff[mask_r]) + 360) % 360
mask_g = (max_rgb == g) & (diff > 0)
h[mask_g] = (60 * ((b[mask_g] - r[mask_g]) / diff[mask_g]) + 120) % 360
mask_b = (max_rgb == b) & (diff > 0)
h[mask_b] = (60 * ((r[mask_b] - g[mask_b]) / diff[mask_b]) + 240) % 360
# Normalize to 0-1 unless degrees mode
if not use_degrees:
h = h / 360.0
# Saturation
s = torch.where(max_rgb > 0, diff / max_rgb, torch.zeros_like(max_rgb))
# Value
v = max_rgb
# Return format
if num_args <= 2:
return torch.stack([h, s, v], dim=-1)
@@ -3193,19 +3193,19 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitHsvToRgbFunc(self, ctx):
num_args = len(ctx.expr())
# Determine mode
if num_args == 1 or num_args == 2:
# Tensor mode
hsv = self._promote_to_tensor((yield ctx.expr(0)))
if hsv.shape[-1] != 3:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: hsv_to_rgb expects tensor with last dim=3, got shape {hsv.shape}")
h = hsv[..., 0]
s = hsv[..., 1]
v = hsv[..., 2]
use_degrees = False
if num_args == 2:
degrees_val = yield ctx.expr(1)
@@ -3215,57 +3215,57 @@ class UnifiedMathVisitor(MathExprVisitor):
h = self._promote_to_tensor((yield ctx.expr(0)))
s = self._promote_to_tensor((yield ctx.expr(1)))
v = self._promote_to_tensor((yield ctx.expr(2)))
use_degrees = False
if num_args == 4:
degrees_val = yield ctx.expr(3)
use_degrees = bool(degrees_val.item() if self._is_tensor(degrees_val) else degrees_val)
# Convert normalized hue to degrees if needed
if not use_degrees:
h = h * 360.0
h = h % 360
s = torch.clamp(s, 0, 1)
v = torch.clamp(v, 0, 1)
# HSV to RGB conversion
c = v * s
x = c * (1 - torch.abs((h / 60) % 2 - 1))
m = v - c
r = torch.zeros_like(h)
g = torch.zeros_like(h)
b = torch.zeros_like(h)
mask0 = (h >= 0) & (h < 60)
r[mask0] = c[mask0]
g[mask0] = x[mask0]
mask1 = (h >= 60) & (h < 120)
r[mask1] = x[mask1]
g[mask1] = c[mask1]
mask2 = (h >= 120) & (h < 180)
g[mask2] = c[mask2]
b[mask2] = x[mask2]
mask3 = (h >= 180) & (h < 240)
g[mask3] = x[mask3]
b[mask3] = c[mask3]
mask4 = (h >= 240) & (h < 300)
r[mask4] = x[mask4]
b[mask4] = c[mask4]
mask5 = (h >= 300) & (h < 360)
r[mask5] = c[mask5]
b[mask5] = x[mask5]
r = r + m
g = g + m
b = b + m
# Return format
if num_args <= 2:
return torch.stack([r, g, b], dim=-1)