4642 lines
178 KiB
Python
4642 lines
178 KiB
Python
import time
|
|
import os
|
|
import torch
|
|
import math
|
|
import inspect
|
|
import torch.nn.functional as F
|
|
from . import optical_flow_utils as ofu
|
|
from .Func import LambdaFunction
|
|
from antlr4 import TerminalNode
|
|
|
|
from .antlr_router import get_antlr_modules
|
|
_, _, MathExprVisitor = get_antlr_modules()
|
|
|
|
from ..helper_functions import generate_dim_variables
|
|
from .noise_utils import NoiseUtils
|
|
import struct
|
|
|
|
from PIL import Image, ImageDraw, ImageFont
|
|
import numpy as np
|
|
|
|
|
|
class ReturnSignal:
|
|
__slots__ = ("value",)
|
|
def __init__(self, value):
|
|
self.value = value
|
|
|
|
class BreakSignal:
|
|
pass
|
|
|
|
class ContinueSignal:
|
|
pass
|
|
|
|
class UnifiedMathVisitor(MathExprVisitor):
|
|
def __init__(self, variables, shape=None, device=None, functions=None, depth=0, state_storage=None):
|
|
|
|
self.variables = variables
|
|
self.spatial_variables = variables.copy()
|
|
self.shape = shape if shape is not None else (1,)
|
|
if device is None:
|
|
self.device = next((v.device for v in variables.values() if isinstance(v, torch.Tensor)), torch.device("cpu"))
|
|
else:
|
|
self.device = device
|
|
self.functions = functions if functions is not None else {}
|
|
self.depth = depth
|
|
self._scope_stack = []
|
|
self._state_storage = state_storage if state_storage is not None else {}
|
|
|
|
def visit(self, tree):
|
|
if tree is None:
|
|
return None
|
|
|
|
gen = tree.accept(self)
|
|
if not inspect.isgenerator(gen):
|
|
return gen
|
|
|
|
stack = [gen]
|
|
last_result = None
|
|
|
|
while stack:
|
|
try:
|
|
# If we're bubbling a signal, we need to check if the parent can handle it.
|
|
if isinstance(last_result, (ReturnSignal, BreakSignal, ContinueSignal)):
|
|
parent_gen = stack[-1]
|
|
func_name = parent_gen.gi_code.co_name
|
|
|
|
is_handler = False
|
|
if isinstance(last_result, (BreakSignal, ContinueSignal)):
|
|
if func_name in ("visitWhileStmt", "visitForStmt"):
|
|
is_handler = True
|
|
elif isinstance(last_result, ReturnSignal):
|
|
if func_name in ("visitCallExp", "visitStart"):
|
|
is_handler = True
|
|
|
|
if not is_handler:
|
|
stack.pop().close()
|
|
continue
|
|
|
|
res = stack[-1].send(last_result)
|
|
|
|
if hasattr(res, 'accept'):
|
|
next_gen = res.accept(self)
|
|
if inspect.isgenerator(next_gen):
|
|
stack.append(next_gen)
|
|
last_result = None
|
|
else:
|
|
last_result = next_gen
|
|
else:
|
|
last_result = res
|
|
|
|
except StopIteration as e:
|
|
stack.pop()
|
|
last_result = e.value
|
|
|
|
if isinstance(last_result, ReturnSignal):
|
|
return last_result.value
|
|
return last_result
|
|
def _is_tensor(self, val):
|
|
return isinstance(val, torch.Tensor) or getattr(val, "is_nested", False)
|
|
|
|
def _is_list(self, val):
|
|
return isinstance(val, (list, tuple))
|
|
|
|
def _promote_to_tensor(self, val,brodcast=False):
|
|
if self._is_tensor(val):
|
|
return val.contiguous()
|
|
if self._is_list(val):
|
|
return torch.tensor(val, device=self.device)
|
|
if brodcast:
|
|
t = list(self.shape)
|
|
t[0]=1
|
|
return torch.full(t,val,device=self.device)
|
|
return torch.tensor(val, device=self.device)
|
|
|
|
def _bin_op(self, a, b, torch_op, scalar_op, ctx):
|
|
"""
|
|
Generic binary operation handler.
|
|
"""
|
|
try:
|
|
if self._is_tensor(a) and a.numel() == 1:
|
|
a = float(a.flatten()[0].item())
|
|
if self._is_tensor(b) and b.numel() == 1:
|
|
b = float(b.flatten()[0].item())
|
|
|
|
# one of them is a list and one is tensor
|
|
if self._is_tensor(a) and self._is_list(b):
|
|
if a.shape[0] == len(b):
|
|
A = torch.split(a, 1)
|
|
results = [self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(A, b)]
|
|
# Ensure all results are tensors
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
results = [self._bin_op(a, x, torch_op, scalar_op, ctx) for x in b]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
if self._is_list(a) and self._is_tensor(b):
|
|
if b.shape[0] == len(a):
|
|
B = torch.split(b, 1)
|
|
results = [self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a, B)]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
results = [self._bin_op(x, b, torch_op, scalar_op, ctx) for x in a]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
|
|
if self._is_list(a) and not self._is_tensor(b):
|
|
if self._is_list(b):
|
|
if len(a) != len(b):
|
|
raise ValueError("List length mismatch")
|
|
return [self._bin_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a, b)]
|
|
return [self._bin_op(x, b, torch_op, scalar_op, ctx) for x in a]
|
|
|
|
if not self._is_tensor(a) and self._is_list(b):
|
|
return [self._bin_op(a, x, torch_op, scalar_op, ctx) for x in b]
|
|
|
|
# Handle tensor operations
|
|
if self._is_tensor(a) or self._is_tensor(b):
|
|
orig_dtype = None
|
|
|
|
# Promote float8 to float16 if needed
|
|
if self._is_tensor(a):
|
|
if hasattr(torch, "float8_e4m3fn") and a.dtype == torch.float8_e4m3fn:
|
|
orig_dtype = a.dtype
|
|
a = a.to(torch.float16)
|
|
elif hasattr(torch, "float8_e5m2") and a.dtype == torch.float8_e5m2:
|
|
orig_dtype = a.dtype
|
|
a = a.to(torch.float16)
|
|
|
|
if self._is_tensor(b):
|
|
if hasattr(torch, "float8_e4m3fn") and b.dtype == torch.float8_e4m3fn:
|
|
orig_dtype = b.dtype if orig_dtype is None else orig_dtype
|
|
b = b.to(torch.float16)
|
|
elif hasattr(torch, "float8_e5m2") and b.dtype == torch.float8_e5m2:
|
|
orig_dtype = b.dtype if orig_dtype is None else orig_dtype
|
|
b = b.to(torch.float16)
|
|
|
|
if torch_op:
|
|
res = torch_op(a, b).contiguous()
|
|
else:
|
|
res = scalar_op(a, b)
|
|
|
|
if orig_dtype is not None and self._is_tensor(res):
|
|
res = res.to(orig_dtype)
|
|
|
|
return res
|
|
|
|
return scalar_op(a, b)
|
|
except (ArithmeticError) as e:
|
|
error_prefix = f"{ctx.start.line}:{ctx.start.column}:"
|
|
raise ArithmeticError(f"{error_prefix} {str(e)}")
|
|
|
|
|
|
def _unary_op(self, a, torch_op, scalar_op):
|
|
if self._is_tensor(a) and a.numel() == 1:
|
|
a = float(a.flatten()[0].item())
|
|
|
|
if self._is_list(a):
|
|
return [self._unary_op(x, torch_op, scalar_op) for x in a]
|
|
if self._is_tensor(a):
|
|
orig_dtype = None
|
|
if hasattr(torch, "float8_e4m3fn") and a.dtype == torch.float8_e4m3fn:
|
|
orig_dtype = a.dtype
|
|
a = a.to(torch.float16)
|
|
elif hasattr(torch, "float8_e5m2") and a.dtype == torch.float8_e5m2:
|
|
orig_dtype = a.dtype
|
|
a = a.to(torch.float16)
|
|
|
|
res = torch_op(a).contiguous() if torch_op else scalar_op(a)
|
|
|
|
if orig_dtype is not None and self._is_tensor(res):
|
|
res = res.to(orig_dtype)
|
|
return res
|
|
|
|
return scalar_op(a)
|
|
|
|
def _reduction_op(self, val, torch_op, list_op):
|
|
if self._is_tensor(val):
|
|
res = torch_op(val)
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
if self._is_list(val):
|
|
return list_op(val)
|
|
return val
|
|
|
|
def _to_int(self, x, ctx, context_name="operation", strict=False):
|
|
"""Convert value to int, handling tensors and nested lists recursively"""
|
|
if self._is_tensor(x):
|
|
if x.numel() == 1:
|
|
return int(x.item())
|
|
elif strict:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {context_name} expects scalar dimensions, got tensor with shape {x.shape}")
|
|
else:
|
|
return x.int()
|
|
elif self._is_list(x):
|
|
if len(x) == 1:
|
|
return self._to_int(x[0], ctx, context_name)
|
|
elif strict:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {context_name} expects scalar dimensions, got list with {len(x)} elements")
|
|
else:
|
|
return [self._to_int(v, ctx, context_name) for v in x]
|
|
else:
|
|
return int(float(x))
|
|
|
|
def _normalize_shape_arg(self, shape_arg, ctx, context_name="random"):
|
|
if isinstance(shape_arg, torch.Size):
|
|
dims = list(shape_arg)
|
|
elif self._is_tensor(shape_arg):
|
|
if shape_arg.numel() == 1:
|
|
return (self._to_int(shape_arg, ctx, context_name),)
|
|
dims = shape_arg.flatten().tolist()
|
|
elif self._is_list(shape_arg):
|
|
dims = list(shape_arg)
|
|
else:
|
|
return (self._to_int(shape_arg, ctx, context_name),)
|
|
|
|
return tuple(self._to_int(d, ctx, context_name) for d in dims)
|
|
|
|
# ========================
|
|
# Visitors
|
|
# ========================
|
|
|
|
def visitNumberExp(self, ctx):
|
|
return float(ctx.NUMBER().getText())
|
|
|
|
def visitConstantExp(self, ctx):
|
|
val = ctx.CONSTANT().getText().lower()
|
|
if val == "pi":
|
|
return math.pi
|
|
if val == "e":
|
|
return math.e
|
|
return 0.0
|
|
|
|
def visitVariableExp(self, ctx):
|
|
var_name = ctx.VARIABLE().getText()
|
|
if var_name == "depth":
|
|
return float(self.depth)
|
|
if var_name in self.variables:
|
|
res = self.variables[var_name]
|
|
return res
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Variable '{var_name}' not found")
|
|
|
|
def visitListExp(self, ctx):
|
|
res = []
|
|
for e in ctx.expr():
|
|
res.append((yield e))
|
|
return res
|
|
|
|
def visitStringExp(self, ctx):
|
|
val = yield ctx.STRING().getText()
|
|
val = val[1:-1].replace('\\n', '\n').replace('\\t', '\t').replace('\\r', '\r').replace('\\\\', '\\').replace('\\"', '"').replace("\\'", "'")
|
|
return val
|
|
|
|
def visitParenExp(self, ctx):
|
|
return (yield ctx.expr())
|
|
|
|
def visitNoneExp(self, ctx):
|
|
return None
|
|
|
|
def visitUnaryPlus(self, ctx):
|
|
return self._unary_op((yield ctx.unaryExpr()), lambda x: x, lambda x: +x)
|
|
|
|
def visitUnaryMinus(self, ctx):
|
|
return self._unary_op((yield ctx.unaryExpr()), torch.neg, lambda x: -x)
|
|
|
|
def visitToIndex(self, ctx):
|
|
return (yield ctx.indexExpr())
|
|
|
|
def visitIndexExp(self, ctx):
|
|
val = (yield ctx.indexExpr())
|
|
raw_index_nodes = ctx.expr()
|
|
|
|
indices = []
|
|
for node in raw_index_nodes:
|
|
idx_val = (yield node)
|
|
|
|
if self._is_tensor(idx_val):
|
|
if idx_val.numel() == 1:
|
|
indices.append(int(idx_val.flatten()[0].item()))
|
|
else:
|
|
# Fancy indexing with tensor
|
|
indices.append(idx_val.long())
|
|
elif self._is_list(idx_val):
|
|
# Fancy indexing with list - convert to tensor
|
|
indices.append(torch.tensor(idx_val, dtype=torch.long, device=self.device))
|
|
else:
|
|
indices.append(int(idx_val))
|
|
|
|
# Use standard PyTorch/list indexing
|
|
if self._is_tensor(val):
|
|
if len(indices) > val.ndim:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Expacted up to {val.ndim} dimensions but got {indices}.")
|
|
|
|
for dim, idx in enumerate(indices):
|
|
if isinstance(idx, int):
|
|
size = val.shape[dim]
|
|
if idx < 0 or idx >= size:
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: Index {idx} out of bounds for dimension {dim} with size {size}"
|
|
)
|
|
else:
|
|
idx_tensor = idx
|
|
if self._is_tensor(idx_tensor):
|
|
if idx_tensor.numel() == 0:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Empty tensor for dimension {dim}")
|
|
if torch.any(idx_tensor < 0) or torch.any(idx_tensor >= val.shape[dim]):
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: Index out of bounds for dimension {dim} with size {val.shape[dim]}"
|
|
)
|
|
|
|
idx_tuple = tuple(indices)
|
|
result = val[idx_tuple]
|
|
if self._is_tensor(result):
|
|
if result.numel() == 1:
|
|
return result.item()
|
|
return result.contiguous()
|
|
return result
|
|
elif isinstance(val, str):
|
|
current = val
|
|
for idx in indices:
|
|
if isinstance(idx, torch.Tensor):
|
|
if idx.numel() != 1:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Too many indecies for string with 1 dimension. Got {idx.numel()}")
|
|
idx = int(idx.flatten()[0].item())
|
|
if idx >= len(current) or idx < 0:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Index {idx} out of bounds for string of length {len(current)}")
|
|
current = current[idx]
|
|
return current
|
|
elif self._is_list(val):
|
|
# Navigate through nested lists
|
|
current = val
|
|
for idx in indices:
|
|
if isinstance(idx, torch.Tensor):
|
|
if idx.numel() != 1:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: List index must be a scalar. Got tensor with length of {idx.numel()}")
|
|
idx = int(idx.item())
|
|
if idx >= len(current) or idx < 0:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Index {idx} out of bounds for list of length {len(current)}")
|
|
current = current[idx]
|
|
return current
|
|
error_prefix = f"{ctx.start.line}:{ctx.start.column}:"
|
|
raise ValueError(f"{error_prefix} Indexing only supported on tensors, lists, and strings (found {type(val).__name__})")
|
|
|
|
def visitToAtom(self, ctx):
|
|
return (yield ctx.atom())
|
|
|
|
def visitTernaryExp(self, ctx):
|
|
condition = yield ctx.compExpr()
|
|
|
|
if self._is_tensor(condition):
|
|
true_val = yield ctx.expr(0)
|
|
false_val = yield ctx.expr(1)
|
|
true_t = self._promote_to_tensor(true_val)
|
|
false_t = self._promote_to_tensor(false_val)
|
|
|
|
if true_t.dtype != false_t.dtype:
|
|
true_t = true_t.float()
|
|
false_t = false_t.float()
|
|
|
|
cond_t = torch.isclose(condition.float(), torch.tensor(0.0, device=self.device)) == False
|
|
return torch.where(cond_t, true_t, false_t).contiguous()
|
|
|
|
if self._is_list(condition):
|
|
res = []
|
|
cache_true = None
|
|
cache_false = None
|
|
for i, c in enumerate(condition):
|
|
if c:
|
|
if cache_true is None:
|
|
cache_true = yield ctx.expr(0)
|
|
res.append(cache_true)
|
|
else:
|
|
if cache_false is None:
|
|
cache_false = yield ctx.expr(1)
|
|
res.append(cache_false)
|
|
return res
|
|
|
|
if condition:
|
|
return (yield ctx.expr(0))
|
|
else:
|
|
return (yield ctx.expr(1))
|
|
|
|
# Binary Ops
|
|
def visitAddExp(self, ctx):
|
|
a = yield ctx.addExpr()
|
|
b = yield ctx.mulExpr()
|
|
|
|
if isinstance(a, str) or isinstance(b, str):
|
|
return str(a) + str(b)
|
|
|
|
return self._bin_op(a, b, torch.add, lambda a, b: a + b, ctx)
|
|
|
|
def visitSubExp(self, ctx):
|
|
a = yield ctx.addExpr()
|
|
b = yield ctx.mulExpr()
|
|
return self._bin_op(a, b, torch.sub, lambda a, b: a - b, ctx)
|
|
|
|
def visitMulExp(self, ctx):
|
|
a = yield ctx.mulExpr()
|
|
b = yield ctx.shiftExpr()
|
|
return self._bin_op(a, b, torch.mul, lambda a, b: a * b, ctx)
|
|
|
|
def visitDivExp(self, ctx):
|
|
a = yield ctx.mulExpr()
|
|
b = yield ctx.shiftExpr()
|
|
return self._bin_op(a, b, torch.div, lambda a, b: a / b, ctx)
|
|
|
|
def visitModExp(self, ctx):
|
|
a = yield ctx.mulExpr()
|
|
b = yield ctx.shiftExpr()
|
|
return self._bin_op(a, b, torch.remainder, lambda a, b: a % b, ctx)
|
|
|
|
def visitPowExp(self, ctx):
|
|
a = yield ctx.unaryExpr()
|
|
b = yield ctx.powExpr()
|
|
return self._bin_op(a, b, torch.pow, lambda a, b: a ** b, ctx)
|
|
|
|
def _bool_op(self, a, b, torch_op, scalar_op, ctx=None):
|
|
if isinstance(a, str) or isinstance(b, str):
|
|
a_str = str(a)
|
|
b_str = str(b)
|
|
result = scalar_op(a_str, b_str)
|
|
return float(result)
|
|
|
|
return self._bin_op(a, b, torch_op, scalar_op, ctx)
|
|
|
|
def visitNeExp(self, ctx):
|
|
a = yield ctx.compExpr()
|
|
b = yield ctx.addExpr()
|
|
return self._bool_op(a, b, torch.ne, lambda a, b: a != b, ctx)
|
|
|
|
def visitEqExp(self, ctx):
|
|
a = yield ctx.compExpr()
|
|
b = yield ctx.addExpr()
|
|
return self._bool_op(a, b, torch.eq, lambda a, b: a == b, ctx)
|
|
|
|
def visitGtExp(self, ctx):
|
|
a = yield ctx.compExpr()
|
|
b = yield ctx.addExpr()
|
|
return self._bool_op(a, b, torch.gt, lambda a, b: a > b, ctx)
|
|
|
|
def visitLtExp(self, ctx):
|
|
a = yield ctx.compExpr()
|
|
b = yield ctx.addExpr()
|
|
return self._bool_op(a, b, torch.lt, lambda a, b: a < b, ctx)
|
|
|
|
def visitGeExp(self, ctx):
|
|
a = yield ctx.compExpr()
|
|
b = yield ctx.addExpr()
|
|
return self._bool_op(a, b, torch.ge, lambda a, b: a >= b, ctx)
|
|
|
|
def visitLeExp(self, ctx):
|
|
a = yield ctx.compExpr()
|
|
b = yield ctx.addExpr()
|
|
return self._bool_op(a, b, torch.le, lambda a, b: a <= b, ctx)
|
|
|
|
def visitToAdd(self, ctx):
|
|
return (yield ctx.addExpr())
|
|
|
|
def visitToMul(self, ctx):
|
|
return (yield ctx.mulExpr())
|
|
|
|
def visitToPow(self, ctx):
|
|
return (yield ctx.powExpr())
|
|
|
|
def visitToUnary(self, ctx):
|
|
return (yield ctx.unaryExpr())
|
|
|
|
# Functions
|
|
def visitTimestampFunc(self, ctx):
|
|
return time.time();
|
|
|
|
def visitSinFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.sin, math.sin)
|
|
|
|
def visitCosFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.cos, math.cos)
|
|
|
|
def visitTanFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.tan, math.tan)
|
|
|
|
def visitAsinFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.asin, math.asin)
|
|
|
|
def visitAcosFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.acos, math.acos)
|
|
|
|
def visitAtanFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.atan, math.atan)
|
|
|
|
def visitSinhFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.sinh, math.sinh)
|
|
|
|
def visitCoshFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.cosh, math.cosh)
|
|
|
|
def visitTanhFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.tanh, math.tanh)
|
|
|
|
def visitAsinhFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.asinh, math.asinh)
|
|
|
|
def visitAcoshFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.acosh, math.acosh)
|
|
|
|
def visitAtanhFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.atanh, math.atanh)
|
|
|
|
def visitAbsFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.abs, abs)
|
|
|
|
def visitAbsExp(self, ctx):
|
|
val = (yield ctx.expr())
|
|
if self._is_list(val):
|
|
return float(torch.linalg.norm(self._promote_to_tensor(val)).item())
|
|
if self._is_tensor(val):
|
|
res = torch.linalg.norm(val)
|
|
if res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
return abs(val)
|
|
|
|
def visitSqrtFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.sqrt, math.sqrt)
|
|
|
|
def visitLnFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.log, math.log)
|
|
|
|
def visitLogFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.log10, math.log10)
|
|
|
|
def visitExpFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.exp, math.exp)
|
|
|
|
def visitFloorFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.floor, math.floor)
|
|
|
|
def visitCeilFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.ceil, math.ceil)
|
|
|
|
def visitRoundFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.round, round)
|
|
|
|
def visitSignFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.sign, lambda x: (1.0 if x > 0 else (-1.0 if x < 0 else 0.0)))
|
|
|
|
def visitFractFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), lambda x: x - torch.floor(x), lambda x: x - math.floor(x))
|
|
|
|
def visitGammaFunc(self, ctx):
|
|
torch_gamma = getattr(torch.special, "gamma", None)
|
|
if torch_gamma is None:
|
|
torch_gamma = lambda x: torch.exp(torch.lgamma(x))
|
|
return self._unary_op((yield ctx.expr()), torch_gamma, math.gamma)
|
|
|
|
def visitSigmoidFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.sigmoid, lambda x: 1.0 / (1.0 + math.exp(-x)))
|
|
|
|
def visitReluFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.relu, lambda x: max(0.0, x))
|
|
|
|
def visitSoftplusFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), F.softplus, lambda x: math.log(1.0 + math.exp(x)))
|
|
|
|
def visitGeluFunc(self, ctx):
|
|
return self._unary_op(
|
|
(yield ctx.expr()), F.gelu, lambda x: 0.5 * x * (1 + math.tanh(math.sqrt(2 / math.pi) * (x + 0.044715 * math.pow(x, 3)))),
|
|
)
|
|
|
|
def visitAnglFunc(self, ctx):
|
|
return self._unary_op((yield ctx.expr()), torch.angle, lambda x: math.atan2(0, x) if x < 0 else 0)
|
|
|
|
def visitPrintFunc(self, ctx):
|
|
val = (yield ctx.expr())
|
|
print(f"{val}")
|
|
return val
|
|
|
|
def visitTNormFunc(self, ctx):
|
|
val = (yield ctx.expr())
|
|
if self._is_tensor(val):
|
|
return F.normalize(val, p=2, dim=-1)
|
|
return 1.0 if val != 0 else 0.0
|
|
|
|
def visitSNormFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
if len(ctx.expr()) > 1:
|
|
dim = self._to_int((yield ctx.expr(1)), ctx, "s_norm dimension")
|
|
if self._is_tensor(val):
|
|
return torch.linalg.norm(val, dim=dim)
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: s_norm with dimension argument only supports tensors")
|
|
if self._is_tensor(val):
|
|
res = torch.linalg.norm(val)
|
|
if res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
return abs(val)
|
|
|
|
# Two-argument functions
|
|
def visitPowFunc(self, ctx):
|
|
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,ctx)
|
|
|
|
def visitTMinFunc(self, ctx):
|
|
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,ctx)
|
|
|
|
def visitStepFunc(self, ctx):
|
|
# step(x, edge) = 1 if x >= edge else 0
|
|
return self._bin_op(
|
|
(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, ctx
|
|
)
|
|
|
|
def visitTopkFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
k = (yield ctx.expr(1))
|
|
|
|
if self._is_tensor(k):
|
|
k_val = int(k.flatten()[0].item())
|
|
else:
|
|
k_val = int(k)
|
|
|
|
if self._is_tensor(val):
|
|
size = val.shape[-1]
|
|
k_val = max(1, min(k_val, size))
|
|
score_val = val.abs() if torch.is_complex(val) else val
|
|
|
|
_, indices = torch.topk(score_val, k=k_val, dim=-1)
|
|
indices = indices.contiguous()
|
|
mask = torch.zeros_like(score_val, dtype=torch.bool)
|
|
mask.scatter_(dim=-1, index=indices, value=True)
|
|
|
|
result = torch.where(mask, val, torch.zeros_like(val))
|
|
return result.contiguous()
|
|
|
|
if self._is_list(val):
|
|
k_val = max(0, min(k_val, len(val)))
|
|
try:
|
|
return sorted(val, reverse=True)[:k_val]
|
|
except:
|
|
return val[:k_val]
|
|
|
|
return val
|
|
|
|
def visitBotkFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
k = (yield ctx.expr(1))
|
|
|
|
if self._is_tensor(k):
|
|
k_val = int(k.flatten()[0].item())
|
|
else:
|
|
k_val = int(k)
|
|
|
|
if self._is_tensor(val):
|
|
size = val.shape[-1]
|
|
k_val = max(1, min(k_val, size))
|
|
score_val = val.abs() if torch.is_complex(val) else val
|
|
|
|
_, indices = torch.topk(score_val, k=k_val, dim=-1, largest=False)
|
|
indices = indices.contiguous()
|
|
mask = torch.zeros_like(score_val, dtype=torch.bool)
|
|
mask.scatter_(dim=-1, index=indices, value=True)
|
|
|
|
result = torch.where(mask, val, torch.zeros_like(val))
|
|
return result.contiguous()
|
|
|
|
if self._is_list(val):
|
|
k_val = max(0, min(k_val, len(val)))
|
|
try:
|
|
return sorted(val)[:k_val]
|
|
except:
|
|
return val[:k_val]
|
|
|
|
return val
|
|
|
|
def visitPinvFunc(self, ctx):
|
|
"""Permutation inverse: if input[i] = j, output[j] = i."""
|
|
val = (yield ctx.expr())
|
|
|
|
if self._is_list(val):
|
|
n = len(val)
|
|
result = [0] * n
|
|
for i, v in enumerate(val):
|
|
idx = int(v)
|
|
if 0 <= idx < n:
|
|
result[idx] = i
|
|
return result
|
|
|
|
if self._is_tensor(val):
|
|
val_list = val.flatten().tolist()
|
|
n = len(val_list)
|
|
result = [0] * n
|
|
for i, v in enumerate(val_list):
|
|
idx = int(v)
|
|
if 0 <= idx < n:
|
|
result[idx] = i
|
|
return torch.tensor(result, device=val.device, dtype=val.dtype).reshape(val.shape)
|
|
|
|
return val
|
|
|
|
def visitGetValueFunc(self, ctx):
|
|
var = yield ctx.expr(0)
|
|
pos_list = yield ctx.expr(1)
|
|
|
|
if not self._is_tensor(var):
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: get_value expects a tensor as first argument")
|
|
|
|
if not self._is_list(pos_list) and not self._is_tensor(pos_list):
|
|
pos_list = [pos_list]
|
|
|
|
if self._is_tensor(pos_list):
|
|
pos_list = pos_list.tolist()
|
|
|
|
if len(pos_list) != var.ndim:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Position list length {len(pos_list)} does not match tensor dimensions {var.ndim}")
|
|
|
|
shape = var.shape
|
|
c_strides = [1] * var.ndim
|
|
if var.ndim > 0:
|
|
for i in range(var.ndim - 2, -1, -1):
|
|
c_strides[i] = c_strides[i+1] * shape[i+1]
|
|
|
|
offset = 0
|
|
for i, p in enumerate(pos_list):
|
|
idx = int(p)
|
|
if idx < 0 or idx >= shape[i]:
|
|
raise ValueError(f"{ctx.VARIABLE().getPayload().line}:{ctx.VARIABLE().getPayload().column}: Index {idx} out of bounds for dimension {i} with size {shape[i]}")
|
|
offset += idx * c_strides[i]
|
|
|
|
return var.contiguous().flatten()[offset]
|
|
|
|
def visitBatchShuffleFunc(self, ctx):
|
|
tsr_val = yield ctx.expr(0)
|
|
idx_val = yield ctx.expr(1)
|
|
|
|
tsr = self._promote_to_tensor(tsr_val)
|
|
|
|
if self._is_tensor(idx_val):
|
|
indices = idx_val.long()
|
|
elif self._is_list(idx_val):
|
|
indices = torch.tensor([int(float(x)) for x in idx_val], dtype=torch.long, device=tsr.device)
|
|
else:
|
|
indices = torch.tensor([int(float(idx_val))], dtype=torch.long, device=tsr.device)
|
|
|
|
# Check bounds
|
|
max_idx = tsr.size(0)
|
|
if torch.any(indices < 0) or torch.any(indices >= max_idx):
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Batch index out of bounds (0-{max_idx-1})")
|
|
|
|
return tsr[indices]
|
|
|
|
def visitArgsortFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0)))
|
|
descending = False
|
|
if ctx.expr(1):
|
|
descending = bool((yield ctx.expr(1)))
|
|
return torch.argsort(val, descending=descending)
|
|
|
|
# Three-argument functions
|
|
def visitClampFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
min_v = (yield ctx.expr(1))
|
|
max_v = (yield ctx.expr(2))
|
|
|
|
if self._is_list(val):
|
|
min_scalar = float(min_v.flatten()[0].item()) if self._is_tensor(min_v) else min_v
|
|
max_scalar = float(max_v.flatten()[0].item()) if self._is_tensor(max_v) else max_v
|
|
return [max(min(x, max_scalar), min_scalar) for x in val]
|
|
if any(self._is_tensor(x) for x in [val, min_v, max_v]):
|
|
return torch.clamp(self._promote_to_tensor(val), self._promote_to_tensor(min_v), self._promote_to_tensor(max_v))
|
|
return max(min(val, max_v), min_v)
|
|
|
|
def visitLerpFunc(self, ctx):
|
|
a = (yield ctx.expr(0))
|
|
b = (yield ctx.expr(1))
|
|
w = (yield ctx.expr(2))
|
|
|
|
if self._is_list(w):
|
|
return [self._lerp_helper(a[i] if self._is_list(a) else a,
|
|
b[i] if self._is_list(b) else b,
|
|
t) for i, t in enumerate(w)]
|
|
|
|
return self._lerp_helper(a, b, w)
|
|
|
|
def _lerp_helper(self, a, b, w):
|
|
if any(self._is_tensor(x) for x in [a, b, w]) or any(self._is_list(x) for x in [a, b, w]):
|
|
return torch.lerp(self._promote_to_tensor(a), self._promote_to_tensor(b), self._promote_to_tensor(w))
|
|
return a*(1-w)+b*w
|
|
|
|
def visitSmoothstepFunc(self, ctx):
|
|
x = (yield ctx.expr(0))
|
|
edge0 = (yield ctx.expr(1))
|
|
edge1 = (yield ctx.expr(2))
|
|
|
|
# t = clamp((x - edge0) / (edge1 - edge0), 0.0, 1.0)
|
|
# return t * t * (3.0 - 2.0 * t)
|
|
|
|
if any(self._is_tensor(v) for v in [x, edge0, edge1]):
|
|
x_t, e0_t, e1_t = self._promote_to_tensor(x), self._promote_to_tensor(edge0), self._promote_to_tensor(edge1)
|
|
t = torch.clamp((x_t - e0_t) / (e1_t - e0_t), 0.0, 1.0)
|
|
return t * t * (3.0 - 2.0 * t)
|
|
|
|
t = max(0.0, min(1.0, (x - edge0) / (edge1 - edge0)))
|
|
return t * t * (3.0 - 2.0 * t)
|
|
|
|
def visitRangeFunc(self, ctx):
|
|
s = (yield ctx.expr(0))
|
|
e = (yield ctx.expr(1))
|
|
st = (yield ctx.expr(2))
|
|
arr = torch.arange(s, e, st, device=self.device, dtype=torch.float32)
|
|
return [float(x) for x in arr.tolist()]
|
|
|
|
def visitSmootherstepFunc(self, ctx):
|
|
x = (yield ctx.expr(0))
|
|
edge0 = (yield ctx.expr(1))
|
|
edge1 = (yield ctx.expr(2))
|
|
|
|
def smoother(t):
|
|
return 6*t**5 - 15*t**4 + 10*t**3
|
|
|
|
if any(self._is_tensor(v) for v in [x, edge0, edge1]):
|
|
x_t, e0_t, e1_t = self._promote_to_tensor(x), self._promote_to_tensor(edge0), self._promote_to_tensor(edge1)
|
|
t = torch.clamp((x_t - e0_t) / (e1_t - e0_t), 0.0, 1.0)
|
|
return smoother(t)
|
|
|
|
t = max(0.0, min(1.0, (x - edge0) / (edge1 - edge0)))
|
|
return smoother(t)
|
|
|
|
def visitCropFunc(self, ctx):
|
|
inp = yield ctx.expr(0)
|
|
pos_list = yield ctx.expr(1)
|
|
size_list = yield ctx.expr(2)
|
|
|
|
def to_int_list(x):
|
|
if self._is_list(x): return [int(v) for v in x]
|
|
if self._is_tensor(x): return x.int().tolist()
|
|
return [int(x)]
|
|
|
|
p_l = to_int_list(pos_list)
|
|
s_l = to_int_list(size_list)
|
|
|
|
# Handle strings
|
|
if isinstance(inp, str):
|
|
start = p_l[0] if p_l else 0
|
|
length = s_l[0] if s_l else len(inp)
|
|
start = max(0, start)
|
|
end = min(len(inp), start + length)
|
|
return inp[start:end]
|
|
|
|
inp = self._promote_to_tensor(inp)
|
|
|
|
if len(p_l) != inp.ndim or len(s_l) != inp.ndim:
|
|
# Basic safety fallback if dims don't match, though robust logic might handle slices properly if we truncate?
|
|
# Let's enforce or just take first N?
|
|
# For robustness, we'll assume user provides correct dims or we raise error?
|
|
given_p = len(p_l)
|
|
given_s = len(s_l)
|
|
if len(p_l) != inp.ndim: raise ValueError(f"{ctx.start.line}:{ctx.start.column}: crop: position dim {given_p} != input dim {inp.ndim}")
|
|
if len(s_l) != inp.ndim: raise ValueError(f"{ctx.start.line}:{ctx.start.column}: crop: size dim {given_s} != input dim {inp.ndim}")
|
|
|
|
out_tensor = torch.zeros(tuple(s_l), dtype=inp.dtype, device=inp.device)
|
|
|
|
slices_in = []
|
|
slices_out = []
|
|
|
|
valid_intersection = True
|
|
|
|
for i in range(inp.ndim):
|
|
start = p_l[i]
|
|
length = s_l[i]
|
|
end = start + length
|
|
|
|
in_start = max(0, start)
|
|
in_end = min(inp.shape[i], end)
|
|
|
|
if in_start >= in_end:
|
|
valid_intersection = False
|
|
break
|
|
|
|
slices_in.append(slice(in_start, in_end))
|
|
|
|
out_start = in_start - start
|
|
out_len = in_end - in_start
|
|
slices_out.append(slice(out_start, out_start + out_len))
|
|
|
|
if valid_intersection:
|
|
out_tensor[tuple(slices_out)] = inp[tuple(slices_in)]
|
|
|
|
return out_tensor
|
|
|
|
def visitCubicEaseFunc(self, ctx):
|
|
a, b, t = (yield ctx.expr(0)), (yield ctx.expr(1)), (yield ctx.expr(2))
|
|
def cubic(v):
|
|
return torch.where(v < 0.5, 4 * v**3, 1 - torch.pow(-2 * v + 2, 3) / 2) if self._is_tensor(v) else \
|
|
(4 * v**3 if v < 0.5 else 1 - math.pow(-2 * v + 2, 3) / 2)
|
|
return self._lerp_helper(a, b, cubic(t))
|
|
|
|
def visitSineEaseFunc(self, ctx):
|
|
a, b, t = (yield ctx.expr(0)), (yield ctx.expr(1)), (yield ctx.expr(2))
|
|
def sine(v):
|
|
return -(torch.cos(math.pi * v) - 1) / 2 if self._is_tensor(v) else -(math.cos(math.pi * v) - 1) / 2
|
|
return self._lerp_helper(a, b, sine(t))
|
|
|
|
def visitElasticEaseFunc(self, ctx):
|
|
a, b, t = (yield ctx.expr(0)), (yield ctx.expr(1)), (yield ctx.expr(2))
|
|
# Specific elastic formula (simplified InOut)
|
|
def elastic(v):
|
|
c4 = (2 * math.pi) / 3
|
|
if self._is_tensor(v):
|
|
return torch.where(v <= 0, 0, torch.where(v >= 1, 1,
|
|
torch.where(v < 0.5, -(torch.pow(2, 20 * v - 10) * torch.sin((20 * v - 11.125) * c4)) / 2,
|
|
(torch.pow(2, -20 * v + 10) * torch.sin((20 * v - 11.125) * c4)) / 2 + 1)))
|
|
if v <= 0: return 0
|
|
if v >= 1: return 1
|
|
if v < 0.5: return -(math.pow(2, 20 * v - 10) * math.sin((20 * v - 11.125) * c4)) / 2
|
|
return (math.pow(2, -20 * v + 10) * math.sin((20 * v - 11.125) * c4)) / 2 + 1
|
|
return self._lerp_helper(a, b, elastic(t))
|
|
|
|
# Helpers for visiting generic exprs
|
|
def visitFunc1Exp(self, ctx):
|
|
res = yield ctx.getChild(0)
|
|
return res
|
|
|
|
def visitFunc2Exp(self, ctx):
|
|
res = yield ctx.getChild(0)
|
|
return res
|
|
|
|
def visitFuncNExp(self, ctx):
|
|
res = yield ctx.getChild(0)
|
|
return res
|
|
|
|
def visitAtomExp(self, ctx):
|
|
res = yield ctx.getChild(0)
|
|
return res
|
|
|
|
def visitExpr(self, ctx):
|
|
res = yield ctx.getChild(0)
|
|
return res
|
|
|
|
def visitSMinFunc(self, ctx):
|
|
vals = []
|
|
for e in ctx.expr():
|
|
vals.append((yield e))
|
|
|
|
if all(not self._is_tensor(x) and not self._is_list(x) for x in vals):
|
|
return min(vals)
|
|
|
|
promoted = [self._promote_to_tensor(x) for x in vals]
|
|
if len(promoted) == 1:
|
|
res = torch.min(promoted[0])
|
|
else:
|
|
res = torch.min(torch.stack(torch.broadcast_tensors(*promoted)))
|
|
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return res.item()
|
|
return res
|
|
|
|
def visitSMaxFunc(self, ctx):
|
|
args = []
|
|
for e in ctx.expr():
|
|
args.append((yield e))
|
|
if len(args) == 1:
|
|
# Check if list or scalar
|
|
if not self._is_tensor(args[0]) and not self._is_list(args[0]):
|
|
return args[0]
|
|
if self._is_list(args[0]):
|
|
return max(args[0]) # max of list
|
|
return torch.max(args[0]).item() # Global max of single tensor
|
|
|
|
# Multiple args
|
|
if all(not self._is_tensor(x) and not self._is_list(x) for x in args):
|
|
return max(args)
|
|
|
|
promoted = [self._promote_to_tensor(x) for x in args]
|
|
if len(promoted) == 1:
|
|
res = torch.max(promoted[0])
|
|
else:
|
|
res = torch.max(torch.stack(torch.broadcast_tensors(*promoted)))
|
|
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
|
|
def _fold_nd(self, tsr, spatial_dims):
|
|
original_shape = tsr.shape
|
|
added_dims = 0
|
|
target_rank = spatial_dims + 2
|
|
while tsr.ndim < target_rank:
|
|
tsr = tsr.unsqueeze(1)
|
|
added_dims += 1
|
|
folded = False
|
|
if tsr.ndim > target_rank:
|
|
fold_count = tsr.ndim - target_rank
|
|
new_batch = 1
|
|
for i in range(fold_count + 1):
|
|
new_batch *= tsr.shape[i]
|
|
tsr = tsr.reshape(new_batch, *tsr.shape[fold_count + 1 :])
|
|
folded = True
|
|
else:
|
|
folded = added_dims > 0
|
|
return tsr, original_shape, added_dims, folded
|
|
|
|
def _unfold_nd(self, tsr, original_shape, added_dims, folded):
|
|
spatial_dims = tsr.ndim - 2
|
|
if folded and added_dims == 0:
|
|
target_fold_rank = len(original_shape) - (spatial_dims + 1)
|
|
fold_dims = original_shape[:target_fold_rank]
|
|
tsr = tsr.reshape(*fold_dims, *tsr.shape[1:])
|
|
for _ in range(added_dims):
|
|
if tsr.ndim > len(original_shape) and tsr.shape[1] == 1:
|
|
tsr = tsr.squeeze(1)
|
|
return tsr
|
|
|
|
def visitPermuteFunc(self, ctx):
|
|
tsr = self._promote_to_tensor((yield ctx.expr(0)))
|
|
dims = (yield ctx.expr(1))
|
|
|
|
# Ensure dims is a list of integers
|
|
if isinstance(dims, torch.Tensor):
|
|
dims = dims.flatten().long().tolist()
|
|
elif isinstance(dims, (list, tuple)):
|
|
dims = [int(0.5 + float(d)) for d in dims] # Round floats to safe ints
|
|
elif isinstance(dims, (int, float)):
|
|
dims = [int(0.5 + float(dims))]
|
|
|
|
return tsr.permute(*dims)
|
|
|
|
def visitReshapeFunc(self, ctx):
|
|
tsr = self._promote_to_tensor((yield ctx.expr(0)))
|
|
new_shape = (yield ctx.expr(1))
|
|
|
|
# Ensure new_shape is a list of integers
|
|
if isinstance(new_shape, torch.Tensor):
|
|
new_shape = new_shape.flatten().long().tolist()
|
|
elif isinstance(new_shape, (list, tuple)):
|
|
result = []
|
|
for d in new_shape:
|
|
# Check if dimension is still a tensor or list (likely wrong variable passed)
|
|
if self._is_tensor(d) and d.numel() > 1:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: reshape expects scalar dimensions, got tensor with shape {d.shape}. Did you mean to pass a shape list instead of data?")
|
|
if self._is_list(d) and len(d) > 1:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: reshape expects scalar dimensions, got list with {len(d)} elements. Did you pass a data variable (like V) instead of a shape?")
|
|
result.append(self._to_int(d, ctx, "reshape", strict=True))
|
|
new_shape = result
|
|
elif isinstance(new_shape, (int, float)):
|
|
new_shape = [int(float(new_shape))]
|
|
|
|
# Validate shape compatibility
|
|
original_numel = tsr.numel()
|
|
target_numel = 1
|
|
for dim in new_shape:
|
|
target_numel *= dim
|
|
|
|
if original_numel != target_numel:
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: Cannot reshape tensor of size {original_numel} "
|
|
f"(shape {list(tsr.shape)}) to shape {new_shape} (size {target_numel}). "
|
|
f"Total elements must match."
|
|
)
|
|
|
|
return tsr.reshape(*new_shape)
|
|
|
|
def visitPrintShapeFunc(self, ctx):
|
|
tsr = (yield ctx.expr())
|
|
if hasattr(tsr, "shape"):
|
|
print(tsr.shape)
|
|
else:
|
|
print(f"Scalar/List: {tsr}")
|
|
return tsr
|
|
|
|
def visitSfftFunc(self, ctx):
|
|
old_vars = self.variables
|
|
self.variables = self.spatial_variables.copy()
|
|
try:
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
dims = tuple(range(val.ndim))
|
|
return torch.fft.fftn(val, dim=dims)
|
|
finally:
|
|
self.variables = old_vars
|
|
|
|
def visitSifftFunc(self, ctx):
|
|
old_vars = self.variables
|
|
self.variables = self.variables.copy()
|
|
device = self.device
|
|
|
|
shape_to_use = self.shape if self.shape else (1, 1, 1, 1)
|
|
if len(ctx.expr()) > 1:
|
|
shape_to_use = (yield ctx.expr(1))
|
|
ndim = len(shape_to_use)
|
|
dim_names = ["x", "y", "z", "w", "v", "u"]
|
|
|
|
k_components = []
|
|
for i in range(ndim):
|
|
dim_idx = ndim - 1 - i
|
|
size_d = shape_to_use[dim_idx]
|
|
values = torch.arange(size_d, dtype=torch.float32, device=device)
|
|
view_shape = [1] * ndim
|
|
view_shape[dim_idx] = size_d
|
|
values = values.view(*view_shape).expand(*shape_to_use)
|
|
|
|
if i < len(dim_names):
|
|
var_name = f"K{dim_names[i]}"
|
|
self.variables[var_name] = values
|
|
self.variables[f"F{dim_names[i]}"] = float(size_d)
|
|
|
|
self.variables[f"K_dim{dim_idx}"] = values
|
|
self.variables[f"F_dim{dim_idx}"] = float(size_d)
|
|
k_components.append(values)
|
|
|
|
k_sq_sum = torch.zeros(shape_to_use, device=device)
|
|
for k_val in k_components:
|
|
k_sq_sum = k_sq_sum + k_val**2
|
|
|
|
self.variables["K"] = torch.sqrt(k_sq_sum)
|
|
self.variables["frequency"] = self.variables["K"]
|
|
|
|
if "Kx" in self.variables:
|
|
self.variables["frequency_count"] = self.variables.get("Fx", 1.0)
|
|
|
|
self.variables = self.variables | generate_dim_variables(k_sq_sum)
|
|
|
|
try:
|
|
val = self._promote_to_tensor((yield ctx.expr(0)))
|
|
dims = tuple(range(val.ndim))
|
|
return torch.fft.ifftn(val, dim=dims).real
|
|
finally:
|
|
self.variables = old_vars
|
|
|
|
def visitSwapFunc(self, ctx):
|
|
tsr = self._promote_to_tensor((yield ctx.expr(0)))
|
|
dim_t = (yield ctx.expr(1))
|
|
idx1_t = (yield ctx.expr(2))
|
|
idx2_t = (yield ctx.expr(3))
|
|
|
|
dim = int(dim_t.flatten()[0].item()) if isinstance(dim_t, torch.Tensor) else int(dim_t)
|
|
i = int(idx1_t.flatten()[0].item()) if isinstance(idx1_t, torch.Tensor) else int(idx1_t)
|
|
j = int(idx2_t.flatten()[0].item()) if isinstance(idx2_t, torch.Tensor) else int(idx2_t)
|
|
|
|
while dim < 0:
|
|
dim += tsr.ndim
|
|
while i < 0:
|
|
i += tsr.shape[dim]
|
|
while j < 0:
|
|
j += tsr.shape[dim]
|
|
|
|
indices = torch.arange(tsr.shape[dim], device=tsr.device)
|
|
val_i = indices[i].clone()
|
|
indices[i] = indices[j]
|
|
indices[j] = val_i
|
|
|
|
return torch.index_select(tsr, dim, indices)
|
|
|
|
def _normalize_coord(self, coord, size):
|
|
if size > 1:
|
|
return (coord / (size - 1)) * 2.0 - 1.0
|
|
return torch.zeros_like(coord)
|
|
|
|
def visitSumFunc(self, ctx):
|
|
return self._reduction_op((yield ctx.expr()), torch.sum, sum)
|
|
|
|
def visitCountFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if self._is_list(val):
|
|
return float(len(val))
|
|
if self._is_tensor(val):
|
|
return val.numel()
|
|
return 1.0
|
|
|
|
def visitMeanFunc(self, ctx):
|
|
return self._reduction_op(
|
|
(yield ctx.expr()), lambda x: torch.mean(x.float()), lambda x: sum(x) / len(x) if x else 0.0
|
|
)
|
|
|
|
def visitStdFunc(self, ctx):
|
|
def list_std(val):
|
|
if len(val) < 2:
|
|
return 0.0
|
|
mean = sum(val) / len(val)
|
|
variance = sum((x - mean) ** 2 for x in val) / (len(val) - 1)
|
|
return math.sqrt(variance)
|
|
|
|
val = (yield ctx.expr())
|
|
if self._is_tensor(val):
|
|
res = torch.std(val.float())
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
elif self._is_list(val):
|
|
return list_std(val)
|
|
else:
|
|
return val
|
|
|
|
def visitVarFunc(self, ctx):
|
|
def list_var(val):
|
|
if len(val) < 2:
|
|
return 0.0
|
|
mean = sum(val) / len(val)
|
|
return sum((x - mean) ** 2 for x in val) / (len(val) - 1)
|
|
|
|
val = (yield ctx.expr())
|
|
if self._is_tensor(val):
|
|
res = torch.var(val.float())
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
elif self._is_list(val):
|
|
return list_var(val)
|
|
else:
|
|
return val
|
|
|
|
def _manual_quantile(self, val, q):
|
|
"""Fallback implementation using sort for when torch.quantile fails on large tensors."""
|
|
val_flat = val.flatten().float()
|
|
sorted_val, _ = torch.sort(val_flat)
|
|
n = len(sorted_val)
|
|
if n == 0:
|
|
return torch.zeros_like(q) if self._is_tensor(q) else 0.0
|
|
|
|
# indices = q * (n - 1)
|
|
indices = q * (n - 1)
|
|
low = torch.floor(indices).long()
|
|
high = torch.ceil(indices).long()
|
|
frac = (indices - low).float()
|
|
|
|
# Ensure bounds
|
|
low = torch.clamp(low, 0, n - 1)
|
|
high = torch.clamp(high, 0, n - 1)
|
|
|
|
res = sorted_val[low] + (sorted_val[high] - sorted_val[low]) * frac
|
|
return res
|
|
|
|
def _quartile_helper(self, val, q):
|
|
|
|
if self._is_tensor(q) and not self._is_tensor(val):
|
|
val = self._promote_to_tensor(val)
|
|
|
|
if self._is_tensor(val):
|
|
if not self._is_tensor(q):
|
|
q = torch.tensor(q, device=self.device).float()
|
|
|
|
try:
|
|
if q.ndim > 1:
|
|
q_flat = q.flatten()
|
|
res = torch.quantile(val.float(), q_flat)
|
|
return res.reshape(q.shape)
|
|
|
|
res = torch.quantile(val.float(), q)
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
except RuntimeError as e:
|
|
# Fallback for "input tensor is too large" or other quantile-specific issues
|
|
if "quantile" in str(e).lower() or "too large" in str(e).lower():
|
|
if q.ndim > 1:
|
|
q_flat = q.flatten()
|
|
res = self._manual_quantile(val, q_flat)
|
|
return res.reshape(q.shape)
|
|
res = self._manual_quantile(val, q)
|
|
if self._is_tensor(res) and res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
raise e
|
|
|
|
if self._is_list(val):
|
|
if not val:
|
|
return 0.0
|
|
sorted_data = sorted(val)
|
|
n = len(val)
|
|
pos = (n - 1) * q
|
|
whole = int(pos)
|
|
frac = pos - whole
|
|
if whole + 1 < n:
|
|
return sorted_data[whole] + (sorted_data[whole + 1] - sorted_data[whole]) * frac
|
|
else:
|
|
return sorted_data[whole]
|
|
return val
|
|
|
|
def visitQuartileFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
k = (yield ctx.expr(1))
|
|
|
|
if self._is_tensor(k):
|
|
# q = k * 0.25. Ensure k is treated as int-like (1,2,3)?
|
|
# Old code did int().
|
|
return self._quartile_helper(val, (k.int().float() * 0.25))
|
|
if self._is_list(k):
|
|
return [self._quartile_helper(val, int(x) * 0.25) for x in k]
|
|
|
|
k_val = int(k)
|
|
q = min(1.0, max(0.0, k_val * 0.25))
|
|
return self._quartile_helper(val, q)
|
|
|
|
def visitPercentileFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
p_raw = (yield ctx.expr(1))
|
|
|
|
if self._is_tensor(p_raw):
|
|
# p is 0-100. q = p / 100
|
|
return self._quartile_helper(val, p_raw.float() / 100.0)
|
|
if self._is_list(p_raw):
|
|
return [self._quartile_helper(val, float(x) / 100.0) for x in p_raw]
|
|
|
|
p = float(p_raw)
|
|
q = max(0.0, min(1.0, p / 100.0))
|
|
return self._quartile_helper(val, q)
|
|
|
|
def visitQuantileFunc(self, ctx):
|
|
val = (yield ctx.expr(0))
|
|
q_raw = (yield ctx.expr(1))
|
|
|
|
if self._is_tensor(q_raw):
|
|
return self._quartile_helper(val, q_raw.float())
|
|
if self._is_list(q_raw):
|
|
return [self._quartile_helper(val, float(x)) for x in q_raw]
|
|
|
|
q = float(q_raw)
|
|
q = max(0.0, min(1.0, q))
|
|
return self._quartile_helper(val, q)
|
|
|
|
def visitDotFunc(self, ctx):
|
|
a = self._promote_to_tensor((yield ctx.expr(0)))
|
|
b = self._promote_to_tensor((yield ctx.expr(1)))
|
|
return float(torch.dot(a.flatten(), b.flatten()).item())
|
|
|
|
def visitMomentFunc(self,ctx):
|
|
x = self._promote_to_tensor((yield ctx.expr(0)))
|
|
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,ctx),k,torch.pow,pow,ctx)).item())/x.numel()
|
|
|
|
def visitSortFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
sorted_val, _ = torch.sort(val)
|
|
return sorted_val
|
|
|
|
def visitCossimFunc(self, ctx):
|
|
a = self._promote_to_tensor((yield ctx.expr(0)))
|
|
b = self._promote_to_tensor((yield ctx.expr(1)))
|
|
|
|
try:
|
|
if a.ndim < 1 or b.ndim < 1:
|
|
raise ValueError("cosine similarity requires tensors with at least 1 dimension")
|
|
return F.cosine_similarity(a.float(), b.float(), dim=-1)
|
|
except RuntimeError as e:
|
|
error_msg = f"{ctx.start.line}:{ctx.start.column}: cossim({a.shape}, {b.shape}): Incompatible shapes for cosine similarity - {str(e)}"
|
|
raise ValueError(error_msg)
|
|
except ValueError as e:
|
|
error_msg = f"{ctx.start.line}:{ctx.start.column}: cossim({a.shape}, {b.shape}): {str(e)}"
|
|
raise ValueError(error_msg)
|
|
|
|
def visitRifeFunc(self, ctx):
|
|
img1 = self._promote_to_tensor((yield ctx.expr(0)))
|
|
img2 = self._promote_to_tensor((yield ctx.expr(1)))
|
|
tiling_size = 0
|
|
iterations = 12
|
|
multi_scale = False
|
|
|
|
if len(ctx.expr()) >= 3:
|
|
tiling_size = float((yield ctx.expr(2)))
|
|
if len(ctx.expr()) >= 4:
|
|
iterations = int((yield ctx.expr(3)))
|
|
if len(ctx.expr()) >= 5:
|
|
multi_scale = bool((yield ctx.expr(4)))
|
|
|
|
return ofu.get_optical_flow(img1, img2, tiling_size, iterations, multi_scale)
|
|
|
|
def visitMotionMaskFunc(self, ctx):
|
|
flow = self._promote_to_tensor((yield ctx.expr()))
|
|
return ofu.calculate_occlusion_mask(flow)
|
|
|
|
def visitFlowToImageFunc(self, ctx):
|
|
flow = self._promote_to_tensor((yield ctx.expr()))
|
|
return ofu.flow_to_image(flow)
|
|
|
|
def visitFlowApplyFunc(self, ctx):
|
|
image = self._promote_to_tensor((yield ctx.expr(0)))
|
|
flow = self._promote_to_tensor((yield ctx.expr(1)))
|
|
return ofu.apply_flow(image, flow)
|
|
|
|
def visitFlipFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0)))
|
|
dims = (yield ctx.expr(1))
|
|
|
|
if self._is_list(dims):
|
|
dims_tuple = tuple(int(x) for x in dims)
|
|
elif self._is_tensor(dims):
|
|
dims_tuple = tuple(dims.long().flatten().tolist())
|
|
else:
|
|
dims_tuple = (int(dims),)
|
|
|
|
return torch.flip(val, dims_tuple)
|
|
|
|
def visitCovFunc(self, ctx):
|
|
x = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
y = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
|
|
x_flat = x.flatten()
|
|
y_flat = y.flatten()
|
|
|
|
if x_flat.numel() != y_flat.numel():
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: x and y must have the same number of elements")
|
|
|
|
n = x_flat.numel()
|
|
if n < 2:
|
|
return torch.tensor(0.0, device=self.device)
|
|
|
|
x_mean = torch.mean(x_flat)
|
|
y_mean = torch.mean(y_flat)
|
|
|
|
sum_sq_diff = torch.sum((x_flat - x_mean) * (y_flat - y_mean)).item()
|
|
return sum_sq_diff / (n - 1)
|
|
|
|
def visitMapFunc(self, ctx):
|
|
tensor = self._promote_to_tensor((yield ctx.expr(0)))
|
|
coords = []
|
|
for i in range(1, len(ctx.expr())):
|
|
coords.append(self._promote_to_tensor((yield ctx.expr(i))))
|
|
num_coords = len(coords)
|
|
|
|
if num_coords == 0:
|
|
return tensor
|
|
if num_coords > 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: map() supports max 3 mapping functions.")
|
|
|
|
if tensor.ndim < num_coords:
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: map() requires input tensor to have at least {num_coords} dimensions "
|
|
f"for {num_coords} coordinate function(s), but got tensor with shape {list(tensor.shape)} ({tensor.ndim} dimension(s)). "
|
|
f"Hint: Use reshape() to add spatial dimensions before mapping."
|
|
)
|
|
|
|
spatial_in_shape = tensor.shape[-num_coords:]
|
|
leading_shape = tensor.shape[:-num_coords]
|
|
|
|
batch_size = 1
|
|
for s in leading_shape:
|
|
batch_size *= s
|
|
|
|
input_view = tensor.reshape(batch_size, 1, *spatial_in_shape)
|
|
norm_coords_list = []
|
|
for i in range(num_coords):
|
|
dim_size = spatial_in_shape[i]
|
|
norm = self._normalize_coord(coords[i], dim_size)
|
|
norm_coords_list.append(norm)
|
|
|
|
grid = torch.stack(norm_coords_list[::-1], dim=-1)
|
|
grid_spatial_shape = grid.shape[:-1]
|
|
|
|
try:
|
|
grid_view = grid.reshape(batch_size, *grid_spatial_shape[-(num_coords):], num_coords)
|
|
except RuntimeError:
|
|
grid_view = grid.expand(batch_size, *([-1] * len(grid_spatial_shape)), -1)
|
|
grid_view = grid_view.reshape(batch_size, *grid_view.shape[-(num_coords + 1) : -1], num_coords)
|
|
|
|
if num_coords == 1:
|
|
input_final = input_view.reshape(batch_size, 1, 1, -1)
|
|
gv = grid_view
|
|
while gv.ndim < 3:
|
|
gv = gv.unsqueeze(1)
|
|
y_zeros = torch.zeros_like(gv[..., :1])
|
|
grid_final = torch.cat([gv, y_zeros], dim=-1).unsqueeze(1) # [B, 1, 1, 2]
|
|
output = F.grid_sample(input_final, grid_final, align_corners=True)
|
|
elif num_coords == 2:
|
|
grid_final = grid_view.reshape(batch_size, *grid_view.shape[-3:-1], 2)
|
|
output = F.grid_sample(input_view, grid_final, align_corners=True)
|
|
else:
|
|
grid_final = grid_view.reshape(batch_size, *grid_view.shape[-4:-1], 3)
|
|
output = F.grid_sample(input_view, grid_final, align_corners=True)
|
|
actual_spatial = grid_view.shape[1:-1]
|
|
final_shape = list(leading_shape) + list(actual_spatial)
|
|
return output.reshape(final_shape)
|
|
|
|
def _parse_conv_args(self, ctx):
|
|
input_raw = (yield ctx.expr(0))
|
|
tensor = self._promote_to_tensor(input_raw)
|
|
num_args = len(ctx.expr())
|
|
if num_args < 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: conv() requires at least 3 arguments")
|
|
|
|
kernel_arg_idx = num_args - 1
|
|
|
|
spatial_dims_count = num_args - 2
|
|
if spatial_dims_count not in [1, 2, 3]:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: conv() supports 1D, 2D, or 3D. Found {spatial_dims_count}")
|
|
|
|
kernel_sizes = []
|
|
for i in range(1, 1 + spatial_dims_count):
|
|
val = (yield ctx.expr(i))
|
|
if isinstance(val, torch.Tensor):
|
|
val = int(val.flatten()[0].item())
|
|
kernel_sizes.append(val)
|
|
|
|
# Prepare context for kernel
|
|
coords = [torch.arange(s, device=self.device).float() - (s // 2) for s in kernel_sizes]
|
|
grid = torch.meshgrid(*coords, indexing="ij")
|
|
|
|
dim_names = ["x", "y", "z"]
|
|
old_vars = self.variables
|
|
self.variables = self.variables.copy()
|
|
for i in range(spatial_dims_count):
|
|
self.variables[f"k{dim_names[i].lower()}"] = grid[i]
|
|
self.variables[f"k{dim_names[i].upper()}"] = grid[i]
|
|
|
|
kernel_var_names = ["kW", "kH", "kD"]
|
|
kernel_var_full = ["kernel_width", "kernel_height", "kernel_depth"]
|
|
for i in range(spatial_dims_count):
|
|
size_val = float(kernel_sizes[i])
|
|
self.variables[kernel_var_names[i]] = size_val
|
|
self.variables[kernel_var_full[i]] = size_val
|
|
|
|
original_shape = self.shape
|
|
self.shape = tuple(kernel_sizes)
|
|
|
|
try:
|
|
kernel_val = self._promote_to_tensor((yield ctx.expr(kernel_arg_idx)))
|
|
finally:
|
|
self.shape = original_shape
|
|
self.variables = old_vars
|
|
|
|
return tensor, kernel_val, kernel_sizes, spatial_dims_count
|
|
|
|
def _apply_conv_internal(self, conv_input, kernel_val, kernel_sizes, spatial_dims_count):
|
|
"""
|
|
Executes the convolution with asymmetric padding support for even kernels.
|
|
Input: [Batch, Channel, Spatial...]
|
|
Kernel: [Spatial...] (to be promoted/repeated)
|
|
kernel_sizes: [W, H, D]
|
|
"""
|
|
in_channels = conv_input.size(1)
|
|
|
|
pads = []
|
|
for k in kernel_sizes:
|
|
p_total = int(k) - 1
|
|
p_low = int(k) // 2
|
|
p_high = p_total - p_low
|
|
pads.extend([p_low, p_high])
|
|
|
|
padded_input = torch.nn.functional.pad(conv_input, tuple(pads), mode="constant", value=0)
|
|
|
|
# Ensure kernel_sizes are integers for tensor operations
|
|
actual_kernel_sizes = [int(k) for k in kernel_sizes[::-1]]
|
|
|
|
if kernel_val.numel() == 1:
|
|
kernel_val = kernel_val.expand(tuple(actual_kernel_sizes))
|
|
elif kernel_val.ndim != spatial_dims_count:
|
|
kernel_val = kernel_val.reshape(tuple(actual_kernel_sizes))
|
|
|
|
final_kernel = kernel_val.unsqueeze(0).unsqueeze(0)
|
|
final_kernel = final_kernel.to(conv_input.dtype)
|
|
final_kernel = final_kernel.repeat(in_channels, 1, *([1] * spatial_dims_count))
|
|
|
|
conv_fn = F.conv1d if spatial_dims_count == 1 else (F.conv2d if spatial_dims_count == 2 else F.conv3d)
|
|
|
|
result = conv_fn(padded_input, final_kernel, padding=0, groups=in_channels)
|
|
return result
|
|
|
|
def visitEzConvFunc(self, ctx):
|
|
tensor, kernel_val, kernel_sizes, spatial_dims_count = (yield from self._parse_conv_args(ctx))
|
|
input_ndim = tensor.ndim
|
|
|
|
is_channels_first = False
|
|
if tensor.ndim == spatial_dims_count + 2:
|
|
c_front = tensor.shape[1]
|
|
c_back = tensor.shape[-1]
|
|
if c_front <= 4 and c_front < c_back:
|
|
is_channels_first = True
|
|
elif c_back <= 4 and c_back < c_front:
|
|
is_channels_first = False
|
|
else:
|
|
is_channels_first = c_back > 4
|
|
|
|
if is_channels_first:
|
|
if spatial_dims_count == 1 and tensor.ndim == 3:
|
|
tensor = tensor.permute(0, 2, 1)
|
|
elif spatial_dims_count == 2 and tensor.ndim == 4:
|
|
tensor = tensor.permute(0, 2, 3, 1)
|
|
elif spatial_dims_count == 3:
|
|
if tensor.ndim == 4:
|
|
tensor = tensor.unsqueeze(-1)
|
|
elif tensor.ndim == 5:
|
|
tensor = tensor.permute(0, 2, 3, 4, 1)
|
|
|
|
if tensor.ndim == spatial_dims_count + 1:
|
|
tensor = tensor.unsqueeze(-1) # Add C=1
|
|
elif tensor.ndim == spatial_dims_count:
|
|
tensor = tensor.unsqueeze(0).unsqueeze(-1) # Add B=1, C=1
|
|
|
|
in_channels = tensor.size(-1)
|
|
batch_end_idx = tensor.ndim - 1 - spatial_dims_count
|
|
batch_shape = tensor.shape[:batch_end_idx]
|
|
spatial_shape = tensor.shape[batch_end_idx:-1]
|
|
channels_shape = (tensor.shape[-1],)
|
|
|
|
total_batch = 1
|
|
for s in batch_shape:
|
|
total_batch *= s
|
|
flat_input = tensor.reshape(total_batch, *spatial_shape, in_channels)
|
|
|
|
permute_order = [0, spatial_dims_count + 1] + list(range(1, spatial_dims_count + 1))
|
|
conv_input = flat_input.permute(*permute_order)
|
|
|
|
result = self._apply_conv_internal(conv_input, kernel_val, kernel_sizes, spatial_dims_count)
|
|
|
|
result_permute = [0] + list(range(2, 2 + spatial_dims_count)) + [1]
|
|
out_flat = result.permute(*result_permute)
|
|
|
|
final_shape = batch_shape + spatial_shape + channels_shape
|
|
out = out_flat.reshape(final_shape)
|
|
|
|
if is_channels_first:
|
|
if spatial_dims_count == 1 and out.ndim == 3:
|
|
out = out.permute(0, 2, 1)
|
|
elif spatial_dims_count == 2 and out.ndim == 4:
|
|
out = out.permute(0, 3, 1, 2)
|
|
elif spatial_dims_count == 3:
|
|
if out.ndim == 5 and out.shape[-1] == 1 and input_ndim == 4:
|
|
out = out.squeeze(-1)
|
|
elif out.ndim == 5:
|
|
out = out.permute(0, 4, 1, 2, 3)
|
|
|
|
return out
|
|
|
|
def visitConvFunc(self, ctx):
|
|
tensor, kernel_val, kernel_sizes, spatial_dims_count = (yield from self._parse_conv_args(ctx))
|
|
|
|
# Expect (Batch..., Channel, Spatial...)
|
|
# spatial_dims_count = 1, 2, or 3
|
|
# Must have at least Channel + Spatial dims
|
|
min_dims = spatial_dims_count + 1
|
|
if tensor.ndim < min_dims:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: convolution() input requires at least Channels + Spatial dimensions. Got shape {tensor.shape} for {spatial_dims_count}D conv.")
|
|
|
|
spatial_shape = tensor.shape[-spatial_dims_count:]
|
|
in_channels = tensor.shape[-(spatial_dims_count + 1)]
|
|
batch_shape = tensor.shape[:-(spatial_dims_count + 1)]
|
|
|
|
total_batch = 1
|
|
for s in batch_shape:
|
|
total_batch *= s
|
|
|
|
conv_input = tensor.reshape(total_batch, in_channels, *spatial_shape)
|
|
|
|
result = self._apply_conv_internal(conv_input, kernel_val, kernel_sizes, spatial_dims_count)
|
|
|
|
return result.reshape(*batch_shape, in_channels, *spatial_shape)
|
|
|
|
def visitAppendFunc(self, ctx):
|
|
a = (yield ctx.expr(0))
|
|
b = (yield ctx.expr(1))
|
|
|
|
if a is None:
|
|
return b
|
|
if b is None:
|
|
return a
|
|
|
|
if self._is_tensor(a) and a.ndim == 0:
|
|
a = a.item()
|
|
if self._is_tensor(b) and b.ndim == 0:
|
|
b = b.item()
|
|
|
|
if self._is_tensor(a) or self._is_tensor(b):
|
|
a = self._promote_to_tensor(a)
|
|
b = self._promote_to_tensor(b)
|
|
|
|
if a.ndim == 0:
|
|
a = a.unsqueeze(0)
|
|
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
|
|
def visitStart(self, ctx):
|
|
count = ctx.getChildCount()
|
|
last_res = None
|
|
|
|
for i in range(count):
|
|
child = ctx.getChild(i)
|
|
if isinstance(child, TerminalNode):
|
|
continue
|
|
|
|
res = yield child
|
|
if res is not None:
|
|
last_res = res
|
|
|
|
return last_res
|
|
|
|
def visitExprStatement(self, ctx):
|
|
res = yield ctx.expr()
|
|
return res
|
|
|
|
def visitVarDefStmt(self, ctx):
|
|
res = yield ctx.varDef()
|
|
return res
|
|
|
|
def visitBlockStatement(self, ctx):
|
|
res = yield ctx.block()
|
|
return res
|
|
|
|
def visitBlock(self, ctx):
|
|
vars_before = set(self.variables.keys())
|
|
try:
|
|
val = None
|
|
for stmt in ctx.stmt():
|
|
val = yield stmt
|
|
return val
|
|
finally:
|
|
for v in set(self.variables.keys()) - vars_before:
|
|
del self.variables[v]
|
|
|
|
def visitIfStatement(self, ctx):
|
|
return (yield ctx.ifStmt())
|
|
|
|
def visitIfStmt(self, ctx):
|
|
cond = yield ctx.expr()
|
|
|
|
def truthy(x):
|
|
if isinstance(x, torch.Tensor):
|
|
return torch.any(x != 0).item()
|
|
if isinstance(x, (list, tuple)):
|
|
return any(x)
|
|
return bool(x)
|
|
|
|
if truthy(cond):
|
|
return (yield ctx.stmt(0))
|
|
elif ctx.stmt(1):
|
|
return (yield ctx.stmt(1))
|
|
return None
|
|
|
|
def visitWhileStmt(self, ctx):
|
|
while True:
|
|
cond = yield ctx.expr()
|
|
is_true = (
|
|
torch.any(cond != 0).item() if isinstance(cond, torch.Tensor)
|
|
else any(cond) if isinstance(cond, (list, tuple))
|
|
else bool(cond)
|
|
)
|
|
|
|
if not is_true:
|
|
break
|
|
|
|
res = yield ctx.stmt()
|
|
|
|
if isinstance(res, ReturnSignal):
|
|
return res
|
|
if isinstance(res, BreakSignal):
|
|
break
|
|
if isinstance(res, ContinueSignal):
|
|
continue
|
|
return None
|
|
|
|
def visitForStmt(self, ctx):
|
|
var_name = ctx.VARIABLE().getText()
|
|
iterable = yield ctx.expr()
|
|
|
|
iterator = []
|
|
if self._is_tensor(iterable):
|
|
if iterable.ndim == 0:
|
|
iterator = [iterable]
|
|
else:
|
|
iterator = iterable
|
|
elif self._is_list(iterable):
|
|
iterator = iterable
|
|
else:
|
|
iterator = [iterable]
|
|
|
|
for val in iterator:
|
|
self.variables[var_name] = val
|
|
res = yield ctx.stmt()
|
|
|
|
if isinstance(res, BreakSignal):
|
|
break
|
|
if isinstance(res, ContinueSignal):
|
|
continue
|
|
return None
|
|
|
|
def visitBreakStmt(self, ctx):
|
|
return BreakSignal()
|
|
|
|
def visitContinueStmt(self, ctx):
|
|
return ContinueSignal()
|
|
|
|
def visitReturnStatement(self, ctx):
|
|
res = yield ctx.returnStmt()
|
|
return res
|
|
|
|
def visitReturnStmt(self, ctx):
|
|
val = (yield ctx.expr()) if ctx.expr() else None
|
|
return ReturnSignal(val)
|
|
|
|
def visitVarDef(self, ctx):
|
|
var_name = ctx.VARIABLE().getText()
|
|
expr_list = ctx.expr()
|
|
|
|
assign_op = "="
|
|
if getattr(ctx, "PLUS_EQ", lambda: None)() is not None: assign_op = "+="
|
|
elif getattr(ctx, "MINUS_EQ", lambda: None)() is not None: assign_op = "-="
|
|
elif getattr(ctx, "MULT_EQ", lambda: None)() is not None: assign_op = "*="
|
|
elif getattr(ctx, "DIV_EQ", lambda: None)() is not None: assign_op = "/="
|
|
elif getattr(ctx, "MOD_EQ", lambda: None)() is not None: assign_op = "%="
|
|
|
|
if not ctx.LBRACKET():
|
|
# Standard assignment: x = value
|
|
val = yield expr_list[0]
|
|
|
|
if assign_op != "=":
|
|
if var_name not in self.variables:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Variable '{var_name}' not defined for compound assignment.")
|
|
existing_val = self.variables[var_name]
|
|
if assign_op == "+=":
|
|
val = self._bin_op(existing_val, val, torch.add, lambda a, b: a + b, ctx)
|
|
elif assign_op == "-=":
|
|
val = self._bin_op(existing_val, val, torch.sub, lambda a, b: a - b, ctx)
|
|
elif assign_op == "*=":
|
|
val = self._bin_op(existing_val, val, torch.mul, lambda a, b: a * b, ctx)
|
|
elif assign_op == "/=":
|
|
val = self._bin_op(existing_val, val, torch.div, lambda a, b: a / b, ctx)
|
|
elif assign_op == "%=":
|
|
val = self._bin_op(existing_val, val, torch.remainder, lambda a, b: a % b, ctx)
|
|
|
|
self.variables[var_name] = val
|
|
return val
|
|
|
|
# Indexed assignment: x[i, j...] = value
|
|
# The last expression is the value to assign
|
|
val_expr = expr_list[-1]
|
|
assigned_val = yield val_expr
|
|
|
|
# Evaluate indices
|
|
indices = []
|
|
for i in range(len(expr_list) - 1):
|
|
indices.append((yield expr_list[i]))
|
|
|
|
if var_name not in self.variables:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Variable '{var_name}' not found for indexed assignment.")
|
|
|
|
target = self.variables[var_name]
|
|
|
|
if self._is_tensor(target):
|
|
# Process indices for PyTorch
|
|
torch_indices = []
|
|
for idx in indices:
|
|
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())
|
|
else:
|
|
torch_indices.append(int(idx))
|
|
|
|
idx_tuple = tuple(torch_indices)
|
|
val_t = self._promote_to_tensor(assigned_val)
|
|
|
|
try:
|
|
# Target slice - used to compute expected shape
|
|
target_slice = target[idx_tuple]
|
|
|
|
if assign_op != "=":
|
|
if assign_op == "+=":
|
|
val_t = self._bin_op(target_slice, val_t, torch.add, lambda a, b: a + b, ctx)
|
|
elif assign_op == "-=":
|
|
val_t = self._bin_op(target_slice, val_t, torch.sub, lambda a, b: a - b, ctx)
|
|
elif assign_op == "*=":
|
|
val_t = self._bin_op(target_slice, val_t, torch.mul, lambda a, b: a * b, ctx)
|
|
elif assign_op == "/=":
|
|
val_t = self._bin_op(target_slice, val_t, torch.div, lambda a, b: a / b, ctx)
|
|
elif assign_op == "%=":
|
|
val_t = self._bin_op(target_slice, val_t, torch.remainder, lambda a, b: a % b, ctx)
|
|
|
|
val_t = self._promote_to_tensor(val_t) # Ensure it's still a tensor
|
|
|
|
# Squeeze leading ones to match target slice rank if it's smaller
|
|
# but target_slice.ndim might be 0 if it's a scalar location.
|
|
while val_t.ndim > target_slice.ndim and val_t.shape[0] == 1:
|
|
val_t = val_t.squeeze(0)
|
|
|
|
target[idx_tuple] = val_t
|
|
return assigned_val if assign_op == "=" else val_t
|
|
except Exception as e:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Indexed assignment to '{var_name}' failed: {str(e)}")
|
|
elif self._is_list(target):
|
|
# Recurse through nested lists if multiple indices provided
|
|
curr = target
|
|
for idx in indices[:-1]:
|
|
curr = curr[int(idx + len(curr) if idx < 0 else idx)]
|
|
last_idx = int(indices[-1])
|
|
real_idx = last_idx + len(curr) if last_idx < 0 else last_idx
|
|
|
|
if assign_op != "=":
|
|
existing_val = curr[real_idx]
|
|
if assign_op == "+=":
|
|
new_val = self._bin_op(existing_val, assigned_val, torch.add, lambda a, b: a + b, ctx)
|
|
elif assign_op == "-=":
|
|
new_val = self._bin_op(existing_val, assigned_val, torch.sub, lambda a, b: a - b, ctx)
|
|
elif assign_op == "*=":
|
|
new_val = self._bin_op(existing_val, assigned_val, torch.mul, lambda a, b: a * b, ctx)
|
|
elif assign_op == "/=":
|
|
new_val = self._bin_op(existing_val, assigned_val, torch.div, lambda a, b: a / b, ctx)
|
|
elif assign_op == "%=":
|
|
new_val = self._bin_op(existing_val, assigned_val, torch.remainder, lambda a, b: a % b, ctx)
|
|
curr[real_idx] = new_val
|
|
return new_val
|
|
else:
|
|
curr[real_idx] = assigned_val
|
|
return assigned_val
|
|
else:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Indexed assignment not supported for {type(target)}")
|
|
|
|
|
|
def visitFunctionDef(self, ctx):
|
|
func_name = ctx.VARIABLE().getText()
|
|
params = []
|
|
if ctx.paramList():
|
|
params = [node.getText() for node in ctx.paramList().VARIABLE()]
|
|
|
|
self.functions[func_name] = {
|
|
"params": params,
|
|
"body": ctx.block() if ctx.block() else ctx.expr()
|
|
}
|
|
return None
|
|
|
|
def visitCallExp(self, ctx):
|
|
func_name = ctx.VARIABLE().getText()
|
|
|
|
# 1) Pokud proměnná existuje a je to lambda uložená v variables, aplikuj ji
|
|
if func_name in self.variables and isinstance(self.variables[func_name], LambdaFunction):
|
|
lam = self.variables[func_name]
|
|
# vyhodnotit argumenty
|
|
args = []
|
|
if ctx.exprList():
|
|
for e in ctx.exprList().expr():
|
|
args.append((yield e))
|
|
|
|
# připrav nový scope na základě uzávěrky
|
|
new_vars = lam.env.copy()
|
|
for i, p in enumerate(lam.params):
|
|
new_vars[p] = args[i] if i < len(args) else None
|
|
|
|
# push/pop scope stejným stylem jako pro pojmenované funkce
|
|
self._scope_stack.append(self.variables)
|
|
self.variables = new_vars
|
|
self.depth += 1
|
|
self.variables["depth"] = float(self.depth)
|
|
try:
|
|
res = yield lam.body
|
|
if isinstance(res, ReturnSignal):
|
|
return res.value
|
|
return res
|
|
finally:
|
|
self.variables = self._scope_stack.pop()
|
|
self.depth -= 1
|
|
|
|
# 2) existující uživatelské funkce (bez změn)
|
|
if func_name in self.functions:
|
|
func_def = self.functions[func_name]
|
|
params = func_def["params"]
|
|
|
|
# Evaluate arguments
|
|
args = []
|
|
if ctx.exprList():
|
|
for e in ctx.exprList().expr():
|
|
args.append((yield e))
|
|
|
|
if len(args) != len(params):
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Function '{func_name}' expects {len(params)} arguments, got {len(args)}"
|
|
)
|
|
|
|
# Create a new scope for function execution
|
|
new_vars = self.variables.copy()
|
|
for param, arg in zip(params, args):
|
|
new_vars[param] = arg
|
|
|
|
# Push current variables to scope stack
|
|
self._scope_stack.append(self.variables)
|
|
|
|
# Update variables and depth
|
|
self.variables = new_vars
|
|
self.depth += 1
|
|
self.variables["depth"] = float(self.depth)
|
|
|
|
try:
|
|
# Visit the body using yield (trampoline will handle it)
|
|
res = yield func_def["body"]
|
|
if isinstance(res, ReturnSignal):
|
|
return res.value
|
|
|
|
return res
|
|
|
|
finally:
|
|
# Restore variables and depth
|
|
self.variables = self._scope_stack.pop()
|
|
self.depth -= 1
|
|
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Unknown function or variable: {func_name}")
|
|
|
|
def visitTextImageFunc(self, ctx):
|
|
"""text_image(text, font, size, [max_width], [weight], [angle], [spacing], [italic], [underline])
|
|
|
|
Renders text to a normalised float32 2D tensor [H, W] using Pillow.
|
|
"""
|
|
|
|
# --- Evaluate arguments ---
|
|
def _to_str(v):
|
|
if isinstance(v, str):
|
|
return v
|
|
if self._is_tensor(v):
|
|
return str(v.item())
|
|
return str(v)
|
|
|
|
def _to_float(v):
|
|
if self._is_tensor(v):
|
|
return float(v.item())
|
|
return float(v)
|
|
|
|
def _to_bool(v):
|
|
if isinstance(v, bool):
|
|
return v
|
|
if self._is_tensor(v):
|
|
v = v.item()
|
|
if isinstance(v, str):
|
|
return v.lower() in ("true", "1", "yes")
|
|
return bool(v)
|
|
|
|
text = _to_str((yield ctx.expr(0)))
|
|
font_name = _to_str((yield ctx.expr(1)))
|
|
size = max(1, self._to_int((yield ctx.expr(2)), ctx,"text_image", strict=True))
|
|
size = max(1, self._to_int((yield ctx.expr(2)), ctx,"text_image", strict=True))
|
|
max_width = self._to_int((yield ctx.expr(3)), ctx,"text_image", strict=True) if len(ctx.expr()) > 3 else 0
|
|
weight_val = self._to_int((yield ctx.expr(4)), ctx,"text_image", strict=True) if len(ctx.expr()) > 4 else 400
|
|
angle = _to_float((yield ctx.expr(5))) if len(ctx.expr()) > 5 else 0.0
|
|
spacing = max(0.1, _to_float((yield ctx.expr(6)))) if len(ctx.expr()) > 6 else 1.0
|
|
is_italic = _to_bool((yield ctx.expr(7))) if len(ctx.expr()) > 7 else False
|
|
has_ul = _to_bool((yield ctx.expr(8))) if len(ctx.expr()) > 8 else False
|
|
|
|
# --- Plně granulární mapování tlouštěk (CSS standard) ---
|
|
if weight_val <= 150:
|
|
weight_tier = "thin"
|
|
elif weight_val <= 250:
|
|
weight_tier = "extralight"
|
|
elif weight_val <= 350:
|
|
weight_tier = "light"
|
|
elif weight_val <= 450:
|
|
weight_tier = "regular"
|
|
elif weight_val <= 550:
|
|
weight_tier = "medium"
|
|
elif weight_val <= 650:
|
|
weight_tier = "semibold"
|
|
elif weight_val <= 750:
|
|
weight_tier = "bold"
|
|
elif weight_val <= 850:
|
|
weight_tier = "extrabold"
|
|
else:
|
|
weight_tier = "black"
|
|
|
|
if is_italic:
|
|
style_key = f"{weight_tier}_italic" if weight_tier != "regular" else "italic"
|
|
else:
|
|
style_key = weight_tier
|
|
|
|
# --- Load font ---
|
|
def _normalise_font_name(value):
|
|
return value.lower().replace(" ", "").replace("-", "").replace("_", "")
|
|
|
|
def _axis_name(axis):
|
|
name = axis.get("name", "")
|
|
if isinstance(name, bytes):
|
|
name = name.decode(errors="ignore")
|
|
return str(name).lower()
|
|
|
|
def _axis_value(axis, value):
|
|
minimum = axis.get("minimum", value)
|
|
maximum = axis.get("maximum", value)
|
|
return max(minimum, min(maximum, value))
|
|
|
|
def _apply_font_variations(fnt, exact_weight, italic):
|
|
if not (
|
|
hasattr(fnt, "get_variation_axes")
|
|
and hasattr(fnt, "set_variation_by_axes")
|
|
):
|
|
return fnt
|
|
|
|
try:
|
|
axes = fnt.get_variation_axes()
|
|
values = []
|
|
changed = False
|
|
|
|
for axis in axes:
|
|
axis_name = _axis_name(axis)
|
|
default = axis.get("default", axis.get("minimum", 0))
|
|
|
|
if axis_name in ("weight", "wght"):
|
|
values.append(_axis_value(axis, exact_weight))
|
|
changed = True
|
|
elif axis_name in ("italic", "ital"):
|
|
values.append(_axis_value(axis, 1 if italic else 0))
|
|
changed = True
|
|
elif axis_name in ("slant", "slnt"):
|
|
target = axis.get("minimum", -12) if italic else 0
|
|
values.append(_axis_value(axis, target))
|
|
changed = True
|
|
else:
|
|
values.append(default)
|
|
|
|
if changed:
|
|
fnt.set_variation_by_axes(values)
|
|
except Exception:
|
|
pass
|
|
|
|
return fnt
|
|
|
|
def _load_font(name, sz, target_style, exact_weight):
|
|
if os.path.exists(name):
|
|
try:
|
|
fnt = ImageFont.truetype(name, sz)
|
|
return _apply_font_variations(
|
|
fnt,
|
|
exact_weight,
|
|
"italic" in target_style,
|
|
)
|
|
except Exception:
|
|
pass
|
|
|
|
style_suffixes = {
|
|
"thin": ["thin", "100", "hairline"],
|
|
"extralight": ["extralight", "ultralight", "200"],
|
|
"light": ["light", "300"],
|
|
"regular": ["regular", "normal", "reg", "400", "standard"],
|
|
"medium": ["medium", "500"],
|
|
"semibold": ["semibold", "demibold", "600"],
|
|
"bold": ["bold", "bd", "700"],
|
|
"extrabold": ["extrabold", "ultrabold", "800"],
|
|
"black": ["black", "heavy", "blk", "900"],
|
|
"italic": ["italic", "it", "i"],
|
|
}.get(target_style.replace("_italic", ""), [""])
|
|
|
|
if "italic" in target_style:
|
|
italic_suffixes = []
|
|
for suffix in style_suffixes:
|
|
italic_suffixes.extend([
|
|
f"{suffix}italic",
|
|
f"{suffix}ital",
|
|
f"{suffix}it",
|
|
f"{suffix}oblique",
|
|
])
|
|
style_suffixes = italic_suffixes + [
|
|
"italic", "ital", "oblique", "it"
|
|
]
|
|
|
|
search_dirs = []
|
|
import platform
|
|
sys_platform = platform.system()
|
|
if sys_platform == "Windows":
|
|
search_dirs = [
|
|
os.path.join(os.environ.get("WINDIR", "C:\\Windows"), "Fonts"),
|
|
os.path.expandvars(r"%LOCALAPPDATA%\Microsoft\Windows\Fonts")
|
|
]
|
|
elif sys_platform == "Darwin":
|
|
search_dirs = ["/Library/Fonts", "/System/Library/Fonts", os.path.expanduser("~/Library/Fonts")]
|
|
else:
|
|
search_dirs = ["/usr/share/fonts", "/usr/local/share/fonts", os.path.expanduser("~/.fonts")]
|
|
|
|
name_clean = _normalise_font_name(name)
|
|
wants_italic = "italic" in target_style
|
|
candidates = []
|
|
|
|
def _font_traits(base_name, sub_name):
|
|
combined = f"{base_name} {sub_name}"
|
|
is_italic_font = any(
|
|
token in combined
|
|
for token in ("italic", "ital", "oblique")
|
|
)
|
|
is_italic_font = is_italic_font or base_name.endswith(
|
|
(f"{name_clean}i", f"{name_clean}bi")
|
|
)
|
|
|
|
weight = 400
|
|
weight_hints = (
|
|
(900, ("black", "heavy", "blk", "900")),
|
|
(800, ("extrabold", "ultrabold", "800")),
|
|
(700, ("bold", "bd", "700")),
|
|
(600, ("semibold", "demibold", "600")),
|
|
(500, ("medium", "500")),
|
|
(300, ("light", "300")),
|
|
(200, ("extralight", "ultralight", "200")),
|
|
(100, ("thin", "hairline", "100")),
|
|
)
|
|
for font_weight, hints in weight_hints:
|
|
if any(hint in combined for hint in hints):
|
|
weight = font_weight
|
|
break
|
|
if base_name.endswith((f"{name_clean}bd", f"{name_clean}bi")):
|
|
weight = 700
|
|
|
|
return is_italic_font, weight
|
|
|
|
def _candidate_score(base_name, sub_name, has_weight_axis):
|
|
is_italic_font, font_weight = _font_traits(base_name, sub_name)
|
|
italic_penalty = 0 if is_italic_font == wants_italic else 10000
|
|
weight_penalty = 0 if has_weight_axis else abs(font_weight - exact_weight)
|
|
exact_name_penalty = 0 if base_name == name_clean else 50
|
|
return italic_penalty + weight_penalty + exact_name_penalty
|
|
|
|
for directory in search_dirs:
|
|
if not os.path.isdir(directory):
|
|
continue
|
|
for root, _, files in os.walk(directory):
|
|
for fname in files:
|
|
if not fname.lower().endswith((".ttf", ".otf")):
|
|
continue
|
|
|
|
full_path = os.path.join(root, fname)
|
|
base_name = _normalise_font_name(os.path.splitext(fname)[0])
|
|
|
|
match_found = False
|
|
has_weight_axis = False
|
|
sub_clean = ""
|
|
|
|
# 1. KONTROLA NÁZVU SOUBORU
|
|
if name_clean in base_name:
|
|
if target_style == "regular" and base_name == name_clean:
|
|
match_found = True
|
|
else:
|
|
match_found = any(suf in base_name for suf in style_suffixes)
|
|
# 2. KONTROLA METADAT UVNITŘ SOUBORU
|
|
if not match_found:
|
|
try:
|
|
test_fnt = ImageFont.truetype(full_path, 10)
|
|
internal_family, internal_sub = test_fnt.getname()
|
|
|
|
family_clean = _normalise_font_name(internal_family)
|
|
sub_clean = _normalise_font_name(internal_sub)
|
|
|
|
if name_clean in family_clean or family_clean in name_clean:
|
|
has_variable_axes = False
|
|
if hasattr(test_fnt, "get_variation_axes"):
|
|
try:
|
|
axes = test_fnt.get_variation_axes()
|
|
axis_names = {
|
|
_axis_name(axis)
|
|
for axis in axes
|
|
}
|
|
has_weight_axis = bool(
|
|
axis_names & {"weight", "wght"}
|
|
)
|
|
has_italic_axis = bool(
|
|
axis_names & {
|
|
"italic",
|
|
"ital",
|
|
"slant",
|
|
"slnt",
|
|
}
|
|
)
|
|
has_variable_axes = (
|
|
has_weight_axis
|
|
or has_italic_axis
|
|
)
|
|
|
|
if wants_italic and has_italic_axis:
|
|
match_found = True
|
|
elif not wants_italic and has_weight_axis:
|
|
match_found = True
|
|
except Exception:
|
|
has_variable_axes = False
|
|
|
|
if not match_found and not has_variable_axes:
|
|
if target_style == "regular":
|
|
match_found = any(
|
|
s in sub_clean
|
|
for s in [
|
|
"regular",
|
|
"normal",
|
|
"standard",
|
|
]
|
|
) or sub_clean == ""
|
|
else:
|
|
match_found = any(
|
|
suf in sub_clean
|
|
for suf in style_suffixes
|
|
) or any(
|
|
suf in family_clean
|
|
for suf in style_suffixes
|
|
)
|
|
except Exception:
|
|
continue
|
|
|
|
if match_found:
|
|
candidates.append((
|
|
_candidate_score(
|
|
base_name,
|
|
sub_clean,
|
|
has_weight_axis,
|
|
),
|
|
full_path,
|
|
))
|
|
|
|
for _, full_path in sorted(candidates, key=lambda item: item[0]):
|
|
try:
|
|
fnt = ImageFont.truetype(full_path, sz)
|
|
return _apply_font_variations(
|
|
fnt,
|
|
exact_weight,
|
|
wants_italic,
|
|
)
|
|
except Exception:
|
|
pass
|
|
|
|
raise FileNotFoundError(
|
|
f"Font '{name}' with style '{target_style}' (weight {exact_weight}) was not found in the system directories."
|
|
)
|
|
|
|
try:
|
|
font = _load_font(font_name, size, style_key, weight_val)
|
|
synthetic_italic = False
|
|
synthetic_bold = False
|
|
except FileNotFoundError:
|
|
try:
|
|
font = _load_font(font_name, size, "regular", 400)
|
|
synthetic_italic = is_italic
|
|
synthetic_bold = (weight_val >= 600)
|
|
except FileNotFoundError:
|
|
raise FileNotFoundError(f"Font '{font_name}' was not found in the system directories (nor its base version).")
|
|
|
|
# --- Word-wrap if max_width is set ---
|
|
def _wrap_text(text_in, fnt, max_w):
|
|
if max_w <= 0:
|
|
return text_in.splitlines()
|
|
lines = []
|
|
for paragraph in text_in.splitlines():
|
|
words = paragraph.split(" ")
|
|
current = ""
|
|
for word in words:
|
|
test = (current + " " + word).strip()
|
|
bbox = fnt.getbbox(test)
|
|
w = bbox[2] - bbox[0]
|
|
if w <= max_w or not current:
|
|
current = test
|
|
else:
|
|
lines.append(current)
|
|
current = word
|
|
lines.append(current)
|
|
return lines
|
|
|
|
lines = _wrap_text(text, font, max_width)
|
|
|
|
# --- Measure canvas size ---
|
|
line_heights = []
|
|
line_widths = []
|
|
for line in lines:
|
|
bbox = font.getbbox(line if line else " ")
|
|
line_widths.append(bbox[2] - bbox[0])
|
|
line_heights.append(bbox[3] - bbox[1])
|
|
|
|
line_h = max(line_heights) if line_heights else size
|
|
line_stride = max(1, int(line_h * spacing))
|
|
|
|
padding_x = int(size * 0.3) if synthetic_italic else 0
|
|
canvas_w = (max(line_widths) if line_widths else size) + padding_x
|
|
canvas_h = line_stride * len(lines)
|
|
|
|
canvas_w = max(canvas_w, 1)
|
|
canvas_h = max(canvas_h, 1)
|
|
|
|
# --- Render text ---
|
|
img = Image.new("L", (canvas_w, canvas_h), color=0)
|
|
draw = ImageDraw.Draw(img)
|
|
y = 0
|
|
|
|
ul_thickness = max(1, int(size * 0.06))
|
|
ul_offset = max(1, int(size * 0.08))
|
|
|
|
for line in lines:
|
|
if not line:
|
|
y += line_stride
|
|
continue
|
|
bbox = font.getbbox(line)
|
|
text_w = bbox[2] - bbox[0]
|
|
text_h = bbox[3] - bbox[1]
|
|
|
|
# Synthetic bold.
|
|
if synthetic_bold:
|
|
thickness_offset = max(1, int(size * 0.02))
|
|
for ox in range(-thickness_offset, thickness_offset + 1):
|
|
for oy in range(-thickness_offset, thickness_offset + 1):
|
|
draw.text((ox, (y - bbox[1]) + oy), line, fill=255, font=font)
|
|
else:
|
|
draw.text((0, y - bbox[1]), line, fill=255, font=font)
|
|
|
|
if has_ul:
|
|
ul_top = y + text_h + ul_offset
|
|
draw.rectangle([0, ul_top, text_w, ul_top + ul_thickness], fill=255)
|
|
|
|
y += line_stride
|
|
|
|
# --- Software Italic (After Text Rendering) ---
|
|
if synthetic_italic:
|
|
img = img.transform(img.size, Image.Transform.AFFINE, (1, 0.25, -0.25*text_h, 0, 1, 0))
|
|
|
|
# --- Rotate if requested ---
|
|
if angle != 0.0:
|
|
img = img.rotate(angle, expand=True, fillcolor=0)
|
|
|
|
# --- Convert to tensor on self.device ---
|
|
arr = np.array(img, dtype=np.float32) / 255.0
|
|
return torch.from_numpy(arr).to(device=self.device)
|
|
|
|
|
|
|
|
def visitNoiseFunc(self,ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 1:
|
|
shape_arg = (yield ctx.expr(1))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "noise")
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return torch.randn(shape_arg, generator=generator, device=self.device)
|
|
|
|
def visitRandFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 1:
|
|
shape_arg = (yield ctx.expr(1))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "rand")
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return torch.rand(shape_arg, generator=generator, device=self.device)
|
|
|
|
def visitExponentialFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 2:
|
|
shape_arg = (yield ctx.expr(2))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "exponential")
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
lambd_val = yield ctx.expr(1)
|
|
lambd = float(lambd_val.item()) if self._is_tensor(lambd_val) else float(lambd_val)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return torch.empty(shape_arg, device=self.device).exponential_(lambd, generator=generator)
|
|
|
|
def visitCauchyFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "cauchy")
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
median_val = yield ctx.expr(1)
|
|
median = float(median_val.item()) if self._is_tensor(median_val) else float(median_val)
|
|
sigma_val = yield ctx.expr(2)
|
|
sigma = float(sigma_val.item()) if self._is_tensor(sigma_val) else float(sigma_val)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return torch.empty(shape_arg, device=self.device).cauchy_(median, sigma, generator=generator)
|
|
|
|
def visitLogNormalFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
mean_val = yield ctx.expr(1)
|
|
mean = float(mean_val.item()) if self._is_tensor(mean_val) else float(mean_val)
|
|
std_val = yield ctx.expr(2)
|
|
std = float(std_val.item()) if self._is_tensor(std_val) else float(std_val)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "log_normal")
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return torch.empty(shape_arg, device=self.device).log_normal_(mean, std, generator=generator)
|
|
|
|
def visitBernoulliFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
p = yield ctx.expr(1)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 2:
|
|
shape_arg = (yield ctx.expr(2))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "bernoulli")
|
|
if self._is_tensor(p):
|
|
return torch.bernoulli(p, generator=generator).to(device=self.device)
|
|
return torch.bernoulli(torch.full(shape_arg, p, device=self.device), generator=generator)
|
|
|
|
def visitPoissonFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
lam = yield ctx.expr(1)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 2:
|
|
shape_arg = (yield ctx.expr(2))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "poisson")
|
|
if self._is_tensor(lam):
|
|
return torch.poisson(lam, generator=generator).to(device=self.device)
|
|
return torch.poisson(torch.full(shape_arg, lam, device=self.device), generator=generator)
|
|
|
|
def visitGammaDistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
shape_val = yield ctx.expr(1)
|
|
shape_param = float(shape_val.item()) if self._is_tensor(shape_val) else float(shape_val)
|
|
scale_val = yield ctx.expr(2)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
shape_arg = self.shape
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "gamma")
|
|
|
|
# Use torch.distributions.Gamma which internally handles the generator properly via torch.manual_seed
|
|
# We set the random state temporarily
|
|
old_state = torch.get_rng_state()
|
|
try:
|
|
torch.manual_seed(seed)
|
|
dist = torch.distributions.Gamma(shape_param, 1.0 / scale)
|
|
result = dist.sample(torch.Size(shape_arg))
|
|
return result.to(device=self.device)
|
|
finally:
|
|
torch.set_rng_state(old_state)
|
|
|
|
def visitBetaDistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
alpha_val = yield ctx.expr(1)
|
|
alpha = float(alpha_val.item()) if self._is_tensor(alpha_val) else float(alpha_val)
|
|
beta_val = yield ctx.expr(2)
|
|
beta = float(beta_val.item()) if self._is_tensor(beta_val) else float(beta_val)
|
|
shape_arg = self.shape
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "beta")
|
|
|
|
old_state = torch.get_rng_state()
|
|
try:
|
|
torch.manual_seed(seed)
|
|
dist = torch.distributions.Beta(alpha, beta)
|
|
result = dist.sample(torch.Size(shape_arg))
|
|
return result.to(device=self.device)
|
|
finally:
|
|
torch.set_rng_state(old_state)
|
|
|
|
def visitLaplaceDistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "laplace")
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
loc_val = yield ctx.expr(1)
|
|
loc = float(loc_val.item()) if self._is_tensor(loc_val) else float(loc_val)
|
|
scale_val = yield ctx.expr(2)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return loc - scale * torch.sign(torch.empty(shape_arg, device=self.device).uniform_(-1, 1, generator=generator)) * torch.log(torch.empty(shape_arg, device=self.device).uniform_(0, 1, generator=generator).clamp(min=1e-10))
|
|
|
|
def visitGumbelDistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "gumbel")
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
loc_val = yield ctx.expr(1)
|
|
loc = float(loc_val.item()) if self._is_tensor(loc_val) else float(loc_val)
|
|
scale_val = yield ctx.expr(2)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
return loc - scale * torch.log(-torch.log(torch.empty(shape_arg, device=self.device).uniform_(0, 1, generator=generator).clamp(min=1e-10)) + 1e-10)
|
|
|
|
def visitWeibullDistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
shape_arg = self.shape;
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
scale_val = yield ctx.expr(1)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
concentration_val = yield ctx.expr(2)
|
|
concentration = float(concentration_val.item()) if self._is_tensor(concentration_val) else float(concentration_val)
|
|
shape_arg = self.shape
|
|
if len(ctx.expr()) > 3:
|
|
shape_arg = (yield ctx.expr(3))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "weibull")
|
|
|
|
# Implement Weibull using generator-aware uniform: scale * (-log(u))^(1/concentration)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
u = torch.rand(shape_arg, generator=generator, device=self.device)
|
|
return scale * torch.pow(-torch.log(u + 1e-10), 1.0 / concentration)
|
|
|
|
def visitChi2DistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
df_val = yield ctx.expr(1)
|
|
df = float(df_val.item()) if self._is_tensor(df_val) else float(df_val)
|
|
shape_arg = self.shape
|
|
if len(ctx.expr()) > 2:
|
|
shape_arg = (yield ctx.expr(2))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "chi2")
|
|
|
|
# Chi-squared is Gamma(df/2, 2)
|
|
old_state = torch.get_rng_state()
|
|
try:
|
|
torch.manual_seed(seed)
|
|
dist = torch.distributions.Gamma(df / 2.0, 0.5)
|
|
result = dist.sample(torch.Size(shape_arg))
|
|
return result.to(device=self.device)
|
|
finally:
|
|
torch.set_rng_state(old_state)
|
|
|
|
def visitStudentTDistFunc(self, ctx):
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
df_val = yield ctx.expr(1)
|
|
df = float(df_val.item()) if self._is_tensor(df_val) else float(df_val)
|
|
shape_arg = self.shape
|
|
if len(ctx.expr()) > 2:
|
|
shape_arg = (yield ctx.expr(2))
|
|
shape_arg = self._normalize_shape_arg(shape_arg, ctx, "student_t")
|
|
|
|
# Student's t using normal and chi-squared: Z / sqrt(V/df) where Z~N(0,1) and V~Chi2(df)
|
|
generator = torch.Generator(device=self.device).manual_seed(seed)
|
|
z = torch.randn(shape_arg, generator=generator, device=self.device)
|
|
|
|
# Generate chi-squared using the same seed + 1 to maintain determinism but different samples
|
|
old_state = torch.get_rng_state()
|
|
try:
|
|
torch.manual_seed(seed + 1)
|
|
dist = torch.distributions.Gamma(df / 2.0, 0.5)
|
|
v = dist.sample(torch.Size(shape_arg))
|
|
v = v.to(device=self.device)
|
|
finally:
|
|
torch.set_rng_state(old_state)
|
|
|
|
return z / torch.sqrt(v / df)
|
|
|
|
def visitNvlFunc(self, ctx):
|
|
v = yield ctx.expr(0)
|
|
v1 = yield ctx.expr(1)
|
|
v2 = yield ctx.expr(2)
|
|
v3 = yield ctx.expr(3)
|
|
return torch.nan_to_num(self._promote_to_tensor(v), v1, v2, v3)
|
|
|
|
def visitAnyFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if self._is_tensor(val): return torch.any(torch.isclose(val, torch.tensor(0.0, device=self.device)) == False).float()
|
|
if self._is_list(val): return float(any(val))
|
|
return float(bool(val))
|
|
|
|
def visitAllFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if self._is_tensor(val): return torch.all(torch.isclose(val, torch.tensor(0.0, device=self.device)) == False).float()
|
|
if self._is_list(val): return float(all(val))
|
|
return float(bool(val))
|
|
|
|
def visitMedianFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if self._is_tensor(val):
|
|
res = torch.median(val.float())
|
|
if res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
if self._is_list(val): return sorted(val)[len(val)//2]
|
|
return val
|
|
|
|
def visitModeFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if self._is_tensor(val):
|
|
res = torch.mode(val.float().flatten()).values
|
|
if res.numel() == 1:
|
|
return float(res.item())
|
|
return res
|
|
if self._is_list(val):
|
|
from collections import Counter
|
|
return Counter(val).most_common(1)[0][0]
|
|
return val
|
|
|
|
def visitCumsumFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
return torch.cumsum(val, dim=0)
|
|
|
|
def visitCumprodFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
return torch.cumprod(val, dim=0)
|
|
|
|
def visitTopkIndFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0)))
|
|
k_val = yield ctx.expr(1)
|
|
k = int(k_val.item()) if self._is_tensor(k_val) else int(k_val)
|
|
return torch.topk(val.flatten(), k=min(k, val.numel()), largest=True).indices
|
|
|
|
def visitBotkIndFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0)))
|
|
k_val = yield ctx.expr(1)
|
|
k = int(k_val.item()) if self._is_tensor(k_val) else int(k_val)
|
|
return torch.topk(val.flatten(), k=min(k, val.numel()), largest=False).indices
|
|
|
|
def _apply_spatial_op(self, tsr, op_fn, original_shape):
|
|
"""
|
|
Helper to handle spatial operations on different layouts.
|
|
Detects [B, H, W, C], [B, C, H, W], [B, H, W], and [H, W, C].
|
|
"""
|
|
ndim = tsr.ndim
|
|
if ndim < 2: return tsr
|
|
|
|
layout = "unknown"
|
|
if ndim == 4:
|
|
# Heuristic: BHWC vs BCHW
|
|
# If last dim is 1, 3, or 4 and much smaller than first/middle dims, likely BHWC
|
|
c_last = original_shape[3]
|
|
if c_last <= 4 and c_last < original_shape[1] and c_last < original_shape[2]:
|
|
tsr = tsr.permute(0, 3, 1, 2)
|
|
layout = "bhwc"
|
|
else:
|
|
# Assume BCHW
|
|
layout = "bchw"
|
|
elif ndim == 3:
|
|
# Heuristic: [B, H, W] (Mask) or [H, W, C] (Image)?
|
|
c_last = original_shape[2]
|
|
if c_last <= 4 and c_last < original_shape[0] and c_last < original_shape[1]:
|
|
# image [H, W, C] -> [1, C, H, W]
|
|
tsr = tsr.permute(2, 0, 1).unsqueeze(0)
|
|
layout = "hwc"
|
|
else:
|
|
# mask [B, H, W] -> [B, 1, H, W]
|
|
tsr = tsr.unsqueeze(1)
|
|
layout = "bhw"
|
|
elif ndim == 2:
|
|
# [H, W] -> [1, 1, H, W]
|
|
tsr = tsr.unsqueeze(0).unsqueeze(0)
|
|
layout = "hw"
|
|
|
|
res = op_fn(tsr)
|
|
|
|
# Restore layout
|
|
if layout == "bhwc":
|
|
return res.permute(0, 2, 3, 1)
|
|
elif layout == "bchw":
|
|
return res
|
|
elif layout == "hwc":
|
|
return res.squeeze(0).permute(1, 2, 0)
|
|
elif layout == "bhw":
|
|
return res.squeeze(1)
|
|
elif layout == "hw":
|
|
return res.squeeze(0).squeeze(0)
|
|
return res
|
|
|
|
def visitEdgeFunc(self, ctx):
|
|
tsr_val = yield ctx.expr(0)
|
|
tsr = self._promote_to_tensor(tsr_val)
|
|
|
|
kernel_size_raw = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
|
|
kernel_size = int(kernel_size_raw.item()) if self._is_tensor(kernel_size_raw) else int(kernel_size_raw)
|
|
|
|
if kernel_size < 3 or (kernel_size % 2) == 0:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: edge kernel_size must be odd and >= 3")
|
|
|
|
original_shape = tsr.shape
|
|
tsr = tsr.float()
|
|
|
|
def build_sobel_kernels(k, device, dtype):
|
|
m = k // 2
|
|
if k == 3:
|
|
d = torch.tensor([-1.0, 0.0, 1.0], device=device, dtype=dtype)
|
|
s = torch.tensor([1.0, 2.0, 1.0], device=device, dtype=dtype)
|
|
else:
|
|
d = torch.arange(-m, m + 1, device=device, dtype=dtype)
|
|
s = torch.tensor([float(math.comb(2 * m, i)) for i in range(2 * m + 1)], device=device, dtype=dtype)
|
|
|
|
kx = torch.outer(s, d)
|
|
ky = torch.outer(d, s)
|
|
|
|
norm_x = torch.sum(torch.abs(kx))
|
|
norm_y = torch.sum(torch.abs(ky))
|
|
if norm_x > 0:
|
|
kx = kx / norm_x
|
|
if norm_y > 0:
|
|
ky = ky / norm_y
|
|
|
|
return kx, ky
|
|
|
|
def sobel_op(x):
|
|
kx, ky = build_sobel_kernels(kernel_size, x.device, x.dtype)
|
|
gx = self._apply_conv_internal(x, kx, [kernel_size, kernel_size], 2)
|
|
gy = self._apply_conv_internal(x, ky, [kernel_size, kernel_size], 2)
|
|
return torch.sqrt(gx**2 + gy**2)
|
|
|
|
return self._apply_spatial_op(tsr, sobel_op, original_shape)
|
|
|
|
def visitGaussianFunc(self, ctx):
|
|
tsr_val = yield ctx.expr(0)
|
|
tsr = self._promote_to_tensor(tsr_val)
|
|
|
|
sigma_val = yield ctx.expr(1)
|
|
sigma = float(sigma_val.item()) if self._is_tensor(sigma_val) else float(sigma_val)
|
|
|
|
if sigma <= 0: return tsr
|
|
original_shape = tsr.shape
|
|
tsr = tsr.float()
|
|
reshap = False
|
|
if len(ctx.expr()) >= 3:
|
|
reshap_val = yield ctx.expr(2)
|
|
reshap = bool(reshap_val.item()) if self._is_tensor(reshap_val) else bool(reshap_val)
|
|
def blur_op(x):
|
|
kernel_size = int(6 * sigma + 1)
|
|
if kernel_size % 2 == 0: kernel_size += 1
|
|
coords = torch.linspace(-kernel_size//2, kernel_size//2, kernel_size, device=x.device)
|
|
kernel = torch.exp(-coords**2 / (2 * sigma**2))
|
|
kernel = kernel / kernel.sum()
|
|
|
|
kh = kernel.view(1, kernel_size)
|
|
x_h = self._apply_conv_internal(x, kh, [kernel_size, 1], 2)
|
|
|
|
kv = kernel.view(kernel_size, 1)
|
|
return self._apply_conv_internal(x_h, kv, [1, kernel_size], 2)
|
|
|
|
return self._apply_spatial_op(tsr, blur_op, original_shape) if reshap else blur_op(tsr)
|
|
|
|
def visitDistFunc(self, ctx):
|
|
x1 = yield ctx.expr(0)
|
|
y1 = yield ctx.expr(1)
|
|
x2 = yield ctx.expr(2)
|
|
y2 = yield ctx.expr(3)
|
|
res_sq = (x2-x1)**2 + (y2-y1)**2
|
|
if self._is_tensor(res_sq):
|
|
return torch.sqrt(res_sq)
|
|
return math.sqrt(res_sq)
|
|
|
|
def visitRemapFunc(self, ctx):
|
|
v = yield ctx.expr(0)
|
|
i_min = yield ctx.expr(1)
|
|
i_max = yield ctx.expr(2)
|
|
o_min = yield ctx.expr(3)
|
|
o_max = yield ctx.expr(4)
|
|
epsilon = 1.0e-10
|
|
denom = (i_max - i_min)
|
|
if self._is_tensor(denom):
|
|
denom = torch.where(denom == 0, torch.fill(denom,epsilon), denom)
|
|
elif self._is_list(denom):
|
|
denom = [epsilon if d == 0 else d for d in denom]
|
|
return [o_min + (vi - i_min) * (o_max - o_min) / di for vi, di in zip(v, denom)]
|
|
elif denom == 0:
|
|
denom = epsilon
|
|
|
|
return o_min + (v - i_min) * (o_max - o_min) / denom
|
|
|
|
def _ensure_dict_storage(self):
|
|
if not isinstance(self._state_storage, dict):
|
|
if not self._state_storage:
|
|
self._state_storage = {}
|
|
else:
|
|
self._state_storage = {i: v for i, v in enumerate(self._state_storage)}
|
|
|
|
def visitPushFunc(self, ctx):
|
|
self._ensure_dict_storage()
|
|
f= yield ctx.expr(0)
|
|
slot = int(f)
|
|
if slot not in self._state_storage:
|
|
self._state_storage[slot] = []
|
|
value = yield ctx.expr(1)
|
|
self._state_storage[slot].append(value)
|
|
return value
|
|
|
|
def visitPopFunc(self, ctx):
|
|
self._ensure_dict_storage()
|
|
slot = int((yield ctx.expr()))
|
|
if slot not in self._state_storage or not self._state_storage[slot]:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Pop from empty slot: {slot}")
|
|
return self._state_storage[slot].pop()
|
|
|
|
def visitClearFunc(self, ctx):
|
|
self._ensure_dict_storage()
|
|
slot = int((yield ctx.expr()))
|
|
if slot in self._state_storage:
|
|
self._state_storage[slot] = []
|
|
return None
|
|
|
|
def visitHasFunc(self, ctx):
|
|
self._ensure_dict_storage()
|
|
slot = int((yield ctx.expr()))
|
|
return float(slot in self._state_storage and bool(self._state_storage[slot]))
|
|
|
|
def visitGetFunc(self, ctx):
|
|
self._ensure_dict_storage()
|
|
slot = int((yield ctx.expr()))
|
|
if slot not in self._state_storage:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Get from empty slot: {slot}")
|
|
storage_list = self._state_storage[slot]
|
|
return storage_list[-1] if storage_list else None
|
|
|
|
def visitBreakExp(self, ctx):
|
|
return BreakSignal()
|
|
|
|
def visitContinueExp(self, ctx):
|
|
return ContinueSignal()
|
|
|
|
def visitEmptyTensorFunc(self, ctx):
|
|
value = (yield ctx.expr(0)) if ctx.expr(0) else 0.0
|
|
type = (yield ctx.expr(1)).dtype if ctx.expr(1) else None
|
|
shape_val = yield ctx.indexExpr()
|
|
|
|
if self._is_list(shape_val):
|
|
shape = [self._to_int(v, ctx, "tensor") for v in shape_val]
|
|
elif self._is_tensor(shape_val):
|
|
shape = [self._to_int(v, ctx, "tensor") for v in shape_val.flatten().tolist()]
|
|
else:
|
|
shape = [self._to_int(shape_val, ctx, "tensor")]
|
|
|
|
return torch.full(shape, value, device=self.device,dtype=type)
|
|
|
|
def visitSoftmaxFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
dim = -1
|
|
if len(ctx.expr()) > 1:
|
|
dim_val = (yield ctx.expr(1))
|
|
dim = self._to_int(dim_val, ctx, "softmax dim", strict=True)
|
|
return F.softmax(val, dim=dim)
|
|
|
|
def visitSoftminFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
dim = -1
|
|
if len(ctx.expr()) > 1:
|
|
dim_val = (yield ctx.expr(1))
|
|
dim = self._to_int(dim_val, ctx, "softmin dim", strict=True)
|
|
return F.softmax(-val, dim=dim)
|
|
|
|
def visitArgminFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
if self._is_tensor(val):
|
|
return torch.argmin(val.flatten())
|
|
if self._is_list(val):
|
|
return float(val.index(min(val)))
|
|
return 0.0
|
|
|
|
def visitArgmaxFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
if self._is_tensor(val):
|
|
return torch.argmax(val.flatten())
|
|
if self._is_list(val):
|
|
return float(val.index(max(val)))
|
|
return 0.0
|
|
|
|
def visitUniqueFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
if self._is_tensor(val):
|
|
unique_vals, _ = torch.unique(val.flatten(), return_counts=False, sorted=True)
|
|
return unique_vals
|
|
if self._is_list(val):
|
|
return sorted(list(set(val)))
|
|
return val
|
|
|
|
def visitFlattenFunc(self, ctx):
|
|
val = (yield ctx.expr())
|
|
if self._is_tensor(val):
|
|
return val.flatten()
|
|
|
|
if self._is_list(val):
|
|
return self._flatten_list(val)
|
|
|
|
return val
|
|
|
|
def _flatten_list(self, lst):
|
|
"""Recursivly flatten list"""
|
|
result = []
|
|
for item in lst:
|
|
if self._is_list(item):
|
|
result.extend(self._flatten_list(item))
|
|
else:
|
|
result.append(item)
|
|
return result
|
|
|
|
def visitCrossFunc(self, ctx):
|
|
a = self._promote_to_tensor((yield ctx.expr(0)))
|
|
b = self._promote_to_tensor((yield ctx.expr(1)))
|
|
|
|
try:
|
|
if a.ndim < 1 or b.ndim < 1:
|
|
raise ValueError("Cross product requires at least 1D tensors")
|
|
if a.shape[-1] != 3 or b.shape[-1] != 3:
|
|
raise ValueError("Cross product requires last dimension size = 3")
|
|
|
|
# Float8 handling
|
|
float8_dtypes = {
|
|
getattr(torch, "float8_e4m3fn", None),
|
|
getattr(torch, "float8_e4m3fnuz", None),
|
|
getattr(torch, "float8_e5m2", None),
|
|
getattr(torch, "float8_e5m2fnuz", None),
|
|
}
|
|
float8_dtypes.discard(None)
|
|
|
|
out_dtype = a.dtype if a.dtype == b.dtype else None
|
|
a_work = a
|
|
b_work = b
|
|
if a.dtype in float8_dtypes or b.dtype in float8_dtypes:
|
|
a_work = a.float()
|
|
b_work = b.float()
|
|
|
|
res = torch.cross(a_work, b_work, dim=-1)
|
|
|
|
if out_dtype in float8_dtypes:
|
|
res = res.to(out_dtype)
|
|
|
|
return res
|
|
except ValueError as e:
|
|
error_msg = f"{ctx.start.line}:{ctx.start.column}: cross({a.shape}, {b.shape}): {str(e)}"
|
|
raise ValueError(error_msg)
|
|
|
|
def visitMatmulFunc(self, ctx):
|
|
a = self._promote_to_tensor((yield ctx.expr(0)))
|
|
b = self._promote_to_tensor((yield ctx.expr(1)))
|
|
|
|
try:
|
|
if a.ndim < 1 or b.ndim < 1:
|
|
raise ValueError("matmul requires tensors with at least 1 dimension")
|
|
return torch.matmul(a, b)
|
|
except RuntimeError as e:
|
|
error_msg = f"{ctx.start.line}:{ctx.start.column}: matmul({a.shape}, {b.shape}): Incompatible shapes for matrix multiplication - {str(e)}"
|
|
raise ValueError(error_msg)
|
|
except ValueError as e:
|
|
error_msg = f"{ctx.start.line}:{ctx.start.column}: matmul({a.shape}, {b.shape}): {str(e)}"
|
|
raise ValueError(error_msg)
|
|
|
|
def visitToShift(self, ctx):
|
|
return (yield ctx.shiftExpr())
|
|
|
|
def visitLShiftExp(self, ctx):
|
|
a = yield ctx.shiftExpr()
|
|
b = yield ctx.powExpr()
|
|
return self._bitwise_op(a, b, torch.bitwise_left_shift, self._scalar_bitwise_lshift,ctx)
|
|
|
|
def visitRShiftExp(self, ctx):
|
|
a = yield ctx.shiftExpr()
|
|
b = yield ctx.powExpr()
|
|
return self._bitwise_op(a, b, torch.bitwise_right_shift, self._scalar_bitwise_rshift,ctx)
|
|
|
|
def visitBitAndFunc(self, ctx):
|
|
a = (yield ctx.expr(0))
|
|
b = (yield ctx.expr(1))
|
|
return self._bitwise_op(a, b, lambda x, y: torch.bitwise_and(x, y), lambda x, y: x & y,ctx)
|
|
|
|
def visitBitXorFunc(self, ctx):
|
|
a = (yield ctx.expr(0))
|
|
b = (yield ctx.expr(1))
|
|
return self._bitwise_op(a, b, lambda x, y: torch.bitwise_xor(x, y), lambda x, y: x ^ y,ctx)
|
|
|
|
def visitBitOrFunc(self, ctx):
|
|
a = (yield ctx.expr(0))
|
|
b = (yield ctx.expr(1))
|
|
return self._bitwise_op(a, b, lambda x, y: torch.bitwise_or(x, y), lambda x, y: x | y,ctx)
|
|
|
|
def visitBitNotFunc(self, ctx):
|
|
v = (yield ctx.expr())
|
|
return self._bitwise_not(v)
|
|
|
|
def visitBitCountFunc(self, ctx):
|
|
v = (yield ctx.expr())
|
|
return self._bitwise_popcount(v)
|
|
|
|
def visitShapeFunc(self, ctx):
|
|
val = (yield ctx.expr())
|
|
if self._is_tensor(val):
|
|
# Return shape as a 1D tensor of integers
|
|
return list(val.shape)
|
|
elif self._is_list(val):
|
|
# Return list length as a single-element tensor
|
|
return [len(val)]
|
|
else:
|
|
# Scalar has shape []
|
|
return []
|
|
|
|
def _bitwise_op(self, a, b, torch_op, scalar_op,ctx):
|
|
"""Binary bitwise operation handler supporting tensors, lists, and scalars."""
|
|
try:
|
|
if self._is_tensor(a) and a.numel() == 1:
|
|
a = int(a.flatten()[0].item())
|
|
if self._is_tensor(b) and b.numel() == 1:
|
|
b = int(b.flatten()[0].item())
|
|
except Exception as e:
|
|
raise ValueError(f"Invalid tensor value for bitwise operation: {e}")
|
|
|
|
# Handle tensor-list combinations
|
|
if self._is_tensor(a) and self._is_list(b):
|
|
if a.shape[0] == len(b):
|
|
A = torch.split(a, 1)
|
|
results = [self._bitwise_op(x, y, torch_op, scalar_op,ctx) for x, y in zip(A, b)]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
results = [self._bitwise_op(a, x, torch_op, scalar_op,ctx) for x in b]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
if self._is_list(a) and self._is_tensor(b):
|
|
if b.shape[0] == len(a):
|
|
B = torch.split(b, 1)
|
|
results = [self._bitwise_op(x, y, torch_op, scalar_op,ctx) for x, y in zip(a, B)]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
results = [self._bitwise_op(x, b, torch_op, scalar_op,ctx) for x in a]
|
|
results = [self._promote_to_tensor(r) if not self._is_tensor(r) else r for r in results]
|
|
return torch.cat([r.unsqueeze(0) if r.ndim == 0 else r for r in results], dim=0)
|
|
|
|
# Handle list-list and list-scalar combinations
|
|
if self._is_list(a) and not self._is_tensor(b):
|
|
if self._is_list(b):
|
|
if len(a) != len(b):
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: List length mismatch in bitwise operation")
|
|
return [self._bitwise_op(x, y, torch_op, scalar_op, ctx) for x, y in zip(a, b)]
|
|
return [self._bitwise_op(x, b, torch_op, scalar_op, ctx) for x in a]
|
|
|
|
if not self._is_tensor(a) and self._is_list(b):
|
|
return [self._bitwise_op(a, x, torch_op, scalar_op,ctx) for x in b]
|
|
|
|
# Handle tensor operations
|
|
if self._is_tensor(a) or self._is_tensor(b):
|
|
if torch_op:
|
|
# View tensors as integers if needed (bitwise ops require integer types)
|
|
original_dtype_a = None
|
|
original_dtype_b = None
|
|
|
|
if self._is_tensor(a):
|
|
original_dtype_a = a.dtype
|
|
if a.dtype not in [torch.int8, torch.int16, torch.int32, torch.int64]:
|
|
# View as integer, don't convert values
|
|
elem_size = a.element_size()
|
|
view_dtype = self._get_bitwise_view_dtype(elem_size)
|
|
a = a.view(view_dtype)
|
|
|
|
if self._is_tensor(b):
|
|
original_dtype_b = b.dtype
|
|
if b.dtype not in [torch.int8, torch.int16, torch.int32, torch.int64]:
|
|
# View as integer, don't convert values
|
|
elem_size = b.element_size()
|
|
view_dtype = self._get_bitwise_view_dtype(elem_size)
|
|
b = b.view(view_dtype)
|
|
|
|
result = torch_op(a, b).contiguous()
|
|
|
|
# View back to original dtype if we viewed a as non-integer
|
|
if original_dtype_a is not None and original_dtype_a not in [torch.int8, torch.int16, torch.int32, torch.int64]:
|
|
result = result.view(original_dtype_a)
|
|
# View back to original dtype if we viewed b as non-integer (and didn't already view from a)
|
|
elif original_dtype_b is not None and original_dtype_b not in [torch.int8, torch.int16, torch.int32, torch.int64]:
|
|
result = result.view(original_dtype_b)
|
|
|
|
return result.contiguous()
|
|
return scalar_op(a, b)
|
|
|
|
return scalar_op(a, b)
|
|
def _bitwise_not(self, v):
|
|
"""Unary bitwise NOT handling for tensors, lists and scalars with support for fp16 and int16."""
|
|
if self._is_tensor(v):
|
|
t = self._promote_to_tensor(v)
|
|
elem_size = t.element_size() if hasattr(t, 'element_size') else 4
|
|
view_dtype = self._get_bitwise_view_dtype(elem_size)
|
|
original_dtype = t.dtype
|
|
|
|
bits = t.view(view_dtype)
|
|
res_bits = torch.bitwise_not(bits)
|
|
return res_bits.view(original_dtype).contiguous()
|
|
|
|
if self._is_list(v):
|
|
return [self._bitwise_not(x) for x in v]
|
|
|
|
# Scalar
|
|
if isinstance(v, int):
|
|
return ~v
|
|
|
|
# For floats or other scalars, operate on bit pattern
|
|
fmt = 'd' if isinstance(v, float) else 'q'
|
|
width = struct.calcsize(fmt) * 8
|
|
bit_fmt = 'Q'
|
|
a_bits = struct.unpack(bit_fmt, struct.pack(fmt, v))[0]
|
|
mask = (1 << width) - 1
|
|
res_bits = (~a_bits) & mask
|
|
try:
|
|
return struct.unpack(fmt, struct.pack(bit_fmt, res_bits))[0]
|
|
except struct.error:
|
|
return int(res_bits)
|
|
|
|
def _get_bitwise_view_dtype(self, elem_size):
|
|
"""Get appropriate integer dtype for bitwise operations based on element size."""
|
|
if elem_size == 1:
|
|
return torch.int8
|
|
elif elem_size == 2:
|
|
return torch.int16
|
|
elif elem_size == 4:
|
|
return torch.int32
|
|
elif elem_size == 8:
|
|
return torch.int64
|
|
else:
|
|
return torch.int32
|
|
|
|
def _bitwise_popcount(self, v):
|
|
"""Count the number of set bits (1s) in the binary representation."""
|
|
if self._is_tensor(v):
|
|
v_t = self._promote_to_tensor(v).flatten().long()
|
|
# Use numpy's bin and count for efficiency
|
|
counts = torch.tensor([bin(int(x) & 0xFFFFFFFFFFFFFFFF).count('1') for x in v_t.tolist()],
|
|
dtype=torch.float32, device=v_t.device)
|
|
if counts.numel() == 1:
|
|
return float(counts.item())
|
|
return counts
|
|
|
|
if self._is_list(v):
|
|
return [self._bitwise_popcount(x) for x in v]
|
|
|
|
# Scalar - count set bits
|
|
v_int = int(v)
|
|
return float(bin(v_int & 0xFFFFFFFFFFFFFFFF).count('1'))
|
|
|
|
def _scalar_bitwise_lshift(self, a, b):
|
|
"""Scalar left shift with bit-pattern preservation for floats."""
|
|
b_int = int(b)
|
|
|
|
# If a is already an int, just do the shift
|
|
if isinstance(a, int):
|
|
return a << b_int
|
|
|
|
# For floats, preserve bit pattern
|
|
if isinstance(a, float):
|
|
fmt = 'd' # double (64-bit)
|
|
bit_fmt = 'Q' # unsigned long long
|
|
a_bits = struct.unpack(bit_fmt, struct.pack(fmt, a))[0]
|
|
result_bits = (a_bits << b_int) & ((1 << 64) - 1) # Mask to 64 bits
|
|
try:
|
|
return struct.unpack(fmt, struct.pack(bit_fmt, result_bits))[0]
|
|
except struct.error:
|
|
return float(result_bits & ((1 << 53) - 1)) # Return mantissa if error
|
|
|
|
# Fallback for other types
|
|
return int(a) << b_int
|
|
|
|
def _scalar_bitwise_rshift(self, a, b):
|
|
"""Scalar right shift with bit-pattern preservation for floats."""
|
|
b_int = int(b)
|
|
|
|
# If a is already an int, just do the shift
|
|
if isinstance(a, int):
|
|
return a >> b_int
|
|
|
|
# For floats, preserve bit pattern
|
|
if isinstance(a, float):
|
|
fmt = 'd' # double (64-bit)
|
|
bit_fmt = 'Q' # unsigned long long
|
|
a_bits = struct.unpack(bit_fmt, struct.pack(fmt, a))[0]
|
|
result_bits = a_bits >> b_int
|
|
try:
|
|
return struct.unpack(fmt, struct.pack(bit_fmt, result_bits))[0]
|
|
except struct.error:
|
|
return float(result_bits)
|
|
|
|
# Fallback for other types
|
|
return int(a) >> b_int
|
|
|
|
def visitPerlinFunc(self, ctx):
|
|
"""perlin(seed, scale, [octaves], [offset], [shape])
|
|
Perlin noise with smooth gradients - supports arbitrary dimensions.
|
|
"""
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
|
|
scale_val = yield ctx.expr(1)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
|
|
octaves = 1
|
|
expr_idx = 2
|
|
if len(ctx.expr()) > expr_idx:
|
|
oct_val = yield ctx.expr(expr_idx)
|
|
octaves = int(oct_val.item()) if self._is_tensor(oct_val) else int(oct_val)
|
|
expr_idx += 1
|
|
|
|
offset = None
|
|
if len(ctx.expr()) > expr_idx:
|
|
offset_val = yield ctx.expr(expr_idx)
|
|
offset = offset_val
|
|
expr_idx += 1
|
|
|
|
# Optional shape parameter
|
|
shape = self.shape
|
|
if len(ctx.expr()) > expr_idx:
|
|
shape_arg = (yield ctx.expr(expr_idx))
|
|
if self._is_tensor(shape_arg):
|
|
shape = tuple(shape_arg.long().flatten().tolist())
|
|
elif self._is_list(shape_arg):
|
|
shape = tuple(int(x) for x in shape_arg)
|
|
else:
|
|
shape = (int(shape_arg),)
|
|
|
|
if len(shape) == 0:
|
|
return torch.tensor(0.0, device=self.device)
|
|
|
|
offset_list = None
|
|
if offset is not None:
|
|
if self._is_tensor(offset):
|
|
offset_list = [float(x) for x in offset.flatten().tolist()]
|
|
elif self._is_list(offset):
|
|
offset_list = [float(x) for x in offset]
|
|
else:
|
|
offset_list = [float(offset)]
|
|
|
|
grids = torch.meshgrid(
|
|
*[
|
|
torch.arange(s, dtype=torch.float32, device=self.device)
|
|
+ (offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0)
|
|
for i, s in enumerate(shape)
|
|
],
|
|
indexing='ij'
|
|
)
|
|
valid_indices = [i for i, s in enumerate(shape) if s > 1]
|
|
|
|
if len(valid_indices) > 0:
|
|
grids_optimized = tuple(grids[i] for i in valid_indices)
|
|
else:
|
|
grids_optimized = grids # Záloha pro případ, že by shape byl např. [1, 1, 1]
|
|
|
|
noise = NoiseUtils.perlin_noise_nd(grids_optimized, scale, seed, self.device)
|
|
|
|
if octaves > 1:
|
|
result = noise
|
|
amplitude = 0.5
|
|
frequency = 2.0
|
|
for octa in range(octaves - 1):
|
|
octave_noise = NoiseUtils.perlin_noise_nd(
|
|
grids_optimized,
|
|
scale / frequency,
|
|
seed + octa + 1,
|
|
self.device
|
|
)
|
|
result = result + octave_noise * amplitude
|
|
amplitude *= 0.5
|
|
frequency *= 2.0
|
|
noise = result / (2 - 2**(-octaves))
|
|
target_shape = grids[0].shape
|
|
noise = noise.view(target_shape)
|
|
return noise * 2 + 0.5
|
|
|
|
def visitCellularFunc(self, ctx):
|
|
"""cellular(seed, scale, [jitter], [offset], [shape])
|
|
Cellular/Voronoi noise - supports arbitrary dimensions.
|
|
"""
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
|
|
scale_val = yield ctx.expr(1)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
|
|
jitter = 0.5
|
|
expr_idx = 2
|
|
if len(ctx.expr()) > expr_idx:
|
|
jitter_val = yield ctx.expr(expr_idx)
|
|
jitter = float(jitter_val.item()) if self._is_tensor(jitter_val) else float(jitter_val)
|
|
jitter = max(0.0, min(1.0, jitter))
|
|
expr_idx += 1
|
|
|
|
offset = None
|
|
if len(ctx.expr()) > expr_idx:
|
|
offset_val = yield ctx.expr(expr_idx)
|
|
offset = offset_val
|
|
expr_idx += 1
|
|
|
|
# Optional shape parameter
|
|
shape = self.shape
|
|
if len(ctx.expr()) > expr_idx:
|
|
shape_arg = (yield ctx.expr(expr_idx))
|
|
if self._is_tensor(shape_arg):
|
|
shape = tuple(shape_arg.long().flatten().tolist())
|
|
elif self._is_list(shape_arg):
|
|
shape = tuple(int(x) for x in shape_arg)
|
|
else:
|
|
shape = (int(shape_arg),)
|
|
|
|
if len(shape) == 0:
|
|
return torch.tensor(0.0, device=self.device)
|
|
|
|
offset_list = None
|
|
if offset is not None:
|
|
if self._is_tensor(offset):
|
|
offset_list = [float(x) for x in offset.flatten().tolist()]
|
|
elif self._is_list(offset):
|
|
offset_list = [float(x) for x in offset]
|
|
else:
|
|
offset_list = [float(offset)]
|
|
|
|
grids = torch.meshgrid(
|
|
*[
|
|
torch.arange(s, dtype=torch.float32, device=self.device)
|
|
+ (offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0)
|
|
for i, s in enumerate(shape)
|
|
],
|
|
indexing='ij'
|
|
)
|
|
|
|
noise = NoiseUtils.cellular_noise_nd(grids, scale, jitter, seed, self.device)
|
|
return noise
|
|
|
|
def visitPlasmaFunc(self, ctx):
|
|
"""plasma(seed, scale, [octaves], [offset], [shape])
|
|
Plasma/Turbulence noise - chaotic high-frequency patterns.
|
|
"""
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
|
|
scale_val = yield ctx.expr(1)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
|
|
octaves = 4 # bylo 1
|
|
expr_idx = 2
|
|
if len(ctx.expr()) > expr_idx:
|
|
oct_val = yield ctx.expr(expr_idx)
|
|
octaves = int(oct_val.item()) if self._is_tensor(oct_val) else int(oct_val)
|
|
expr_idx += 1
|
|
|
|
offset = None
|
|
if len(ctx.expr()) > expr_idx:
|
|
offset_val = yield ctx.expr(expr_idx)
|
|
offset = offset_val
|
|
expr_idx += 1
|
|
|
|
shape = self.shape
|
|
if len(ctx.expr()) > expr_idx:
|
|
shape_arg = (yield ctx.expr(expr_idx))
|
|
if self._is_tensor(shape_arg):
|
|
shape = tuple(shape_arg.long().flatten().tolist())
|
|
elif self._is_list(shape_arg):
|
|
shape = tuple(int(x) for x in shape_arg)
|
|
else:
|
|
shape = (int(shape_arg),)
|
|
|
|
if len(shape) == 0:
|
|
return torch.tensor(0.0, device=self.device)
|
|
|
|
offset_list = None
|
|
if offset is not None:
|
|
if self._is_tensor(offset):
|
|
offset_list = [float(x) for x in offset.flatten().tolist()]
|
|
elif self._is_list(offset):
|
|
offset_list = [float(x) for x in offset]
|
|
else:
|
|
offset_list = [float(offset)]
|
|
|
|
grids = torch.meshgrid(
|
|
*[
|
|
torch.arange(s, dtype=torch.float32, device=self.device)
|
|
+ (offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0)
|
|
for i, s in enumerate(shape)
|
|
],
|
|
indexing='ij'
|
|
)
|
|
valid_indices = [i for i, s in enumerate(shape) if s > 1]
|
|
|
|
if len(valid_indices) > 0:
|
|
grids_optimized = tuple(grids[i] for i in valid_indices)
|
|
else:
|
|
grids_optimized = grids
|
|
noise = NoiseUtils.plasma_noise_nd(
|
|
grids_optimized,
|
|
scale=scale,
|
|
seed=seed,
|
|
device=self.device,
|
|
octaves=octaves
|
|
)
|
|
target_shape = grids[0].shape
|
|
noise = noise.view(target_shape)
|
|
return noise
|
|
|
|
def visitPadFunc(self,ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr(0)))
|
|
pad_val = yield ctx.expr(1)
|
|
if self._is_tensor(pad_val):
|
|
pad = [int(x) for x in pad_val.flatten().tolist()]
|
|
elif self._is_list(pad_val):
|
|
pad = [int(x) for x in pad_val]
|
|
else:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Pad amount must be a list or tensor.")
|
|
|
|
if len(pad) % 2 != 0:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Pad amount list must have an even number of elements.")
|
|
|
|
reversed_pad = []
|
|
for i in range(len(pad) - 1, 0, -2):
|
|
reversed_pad.extend([pad[i-1], pad[i]])
|
|
return F.pad(val, reversed_pad)
|
|
|
|
def _align_mask_to_overlay(self, mask, overlay_shape):
|
|
if mask.shape == overlay_shape:
|
|
return mask
|
|
if mask.ndim == 0:
|
|
return mask.expand(overlay_shape)
|
|
|
|
# Try left-to-right alignment
|
|
mask_shape = list(mask.shape)
|
|
new_shape = []
|
|
mask_idx = 0
|
|
for dim_size in overlay_shape:
|
|
if mask_idx < len(mask_shape) and mask_shape[mask_idx] == dim_size:
|
|
new_shape.append(dim_size)
|
|
mask_idx += 1
|
|
elif mask_idx < len(mask_shape) and mask_shape[mask_idx] == 1:
|
|
new_shape.append(1)
|
|
mask_idx += 1
|
|
else:
|
|
new_shape.append(1)
|
|
if mask_idx == len(mask_shape):
|
|
return mask.view(new_shape).expand(overlay_shape)
|
|
|
|
# Try right-to-left alignment
|
|
new_shape = []
|
|
mask_idx = len(mask_shape) - 1
|
|
for dim_size in reversed(overlay_shape):
|
|
if mask_idx >= 0 and mask_shape[mask_idx] == dim_size:
|
|
new_shape.insert(0, dim_size)
|
|
mask_idx -= 1
|
|
elif mask_idx >= 0 and mask_shape[mask_idx] == 1:
|
|
new_shape.insert(0, 1)
|
|
mask_idx -= 1
|
|
else:
|
|
new_shape.insert(0, 1)
|
|
if mask_idx < 0:
|
|
return mask.view(new_shape).expand(overlay_shape)
|
|
|
|
# Fallback: try broadcast_to or expand
|
|
try:
|
|
return torch.broadcast_to(mask, overlay_shape)
|
|
except Exception:
|
|
try:
|
|
return mask.expand(overlay_shape)
|
|
except Exception:
|
|
return mask
|
|
|
|
def visitOverlayFunc(self, ctx):
|
|
base = yield ctx.expr(0)
|
|
overlay = yield ctx.expr(1)
|
|
offset_raw = yield ctx.expr(2)
|
|
|
|
if len(ctx.expr()) > 3:
|
|
opacity = yield ctx.expr(3)
|
|
else:
|
|
opacity = 1.0
|
|
|
|
opacity_is_tensor = self._is_tensor(opacity) and opacity.numel() > 1
|
|
|
|
if not opacity_is_tensor:
|
|
opacity_f = float(opacity.item()) if self._is_tensor(opacity) else float(opacity)
|
|
if opacity_f <= 0.0:
|
|
return base
|
|
else:
|
|
opacity_f = 1.0
|
|
|
|
if isinstance(base, str):
|
|
if not isinstance(overlay, str):
|
|
overlay = str(overlay)
|
|
|
|
offset = int(offset_raw) if not self._is_tensor(offset_raw) else int(offset_raw.item())
|
|
if offset >= len(base):
|
|
return base
|
|
|
|
if offset < 0:
|
|
overlay = overlay[-offset:]
|
|
if opacity_is_tensor:
|
|
opacity = opacity[-offset:]
|
|
offset = 0
|
|
|
|
end = min(len(base), offset + len(overlay))
|
|
overlay_len = end - offset
|
|
|
|
if not opacity_is_tensor and opacity_f >= 1.0:
|
|
return base[:offset] + overlay[:overlay_len] + base[end:]
|
|
|
|
mixed = list(base)
|
|
if opacity_is_tensor:
|
|
opacity_list = opacity.detach().cpu().flatten().tolist()
|
|
for i in range(overlay_len):
|
|
a = ord(base[offset + i])
|
|
c = ord(overlay[i])
|
|
op_val = opacity_list[i] if (opacity_is_tensor and i < len(opacity_list)) else (opacity_f if not opacity_is_tensor else 1.0)
|
|
avg = round(a * (1 - op_val) + c * op_val)
|
|
mixed[offset + i] = chr(avg)
|
|
return ''.join(mixed)
|
|
|
|
if self._is_list(base):
|
|
if not self._is_list(overlay):
|
|
overlay = [overlay]
|
|
|
|
offset = int(offset_raw) if not self._is_tensor(offset_raw) else int(offset_raw.item())
|
|
if offset >= len(base):
|
|
return base
|
|
|
|
if offset < 0:
|
|
overlay = overlay[-offset:]
|
|
if opacity_is_tensor:
|
|
opacity = opacity[-offset:]
|
|
offset = 0
|
|
|
|
result = list(base)
|
|
end = min(len(base), offset + len(overlay))
|
|
if opacity_is_tensor:
|
|
opacity_list = opacity.detach().cpu().flatten().tolist()
|
|
for i, val in enumerate(overlay[:end - offset]):
|
|
op_val = opacity_list[i] if (opacity_is_tensor and i < len(opacity_list)) else (opacity_f if not opacity_is_tensor else 1.0)
|
|
if op_val >= 1.0:
|
|
result[offset + i] = val
|
|
elif op_val <= 0.0:
|
|
pass
|
|
else:
|
|
result[offset + i] = result[offset + i] * (1.0 - op_val) + val * op_val
|
|
return result
|
|
|
|
# Tensor path
|
|
base = self._promote_to_tensor(base)
|
|
overlay = self._promote_to_tensor(overlay)
|
|
mask = self._promote_to_tensor(opacity)
|
|
|
|
if opacity_is_tensor:
|
|
mask = mask.to(dtype=overlay.dtype, device=overlay.device)
|
|
mask = self._align_mask_to_overlay(mask, overlay.shape)
|
|
|
|
if self._is_tensor(offset_raw):
|
|
offset = [int(x) for x in offset_raw.flatten().tolist()]
|
|
elif self._is_list(offset_raw):
|
|
offset = [int(x) for x in offset_raw]
|
|
else:
|
|
offset = [int(offset_raw)]
|
|
|
|
if len(offset) != base.ndim:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Offset dimensions {len(offset)} must match base dimensions {base.ndim}")
|
|
|
|
crop_slices = []
|
|
paste_slices = []
|
|
|
|
for i in range(base.ndim):
|
|
off = offset[i]
|
|
bsz = base.shape[i]
|
|
osz = overlay.shape[i]
|
|
|
|
if off >= bsz or (off + osz) <= 0:
|
|
return base
|
|
|
|
paste_start = max(off, 0)
|
|
paste_end = min(off + osz, bsz)
|
|
|
|
crop_start = max(-off, 0)
|
|
crop_end = crop_start + (paste_end - paste_start)
|
|
|
|
crop_slices.append(slice(crop_start, crop_end))
|
|
paste_slices.append(slice(paste_start, paste_end))
|
|
|
|
cropped_overlay = overlay[tuple(crop_slices)]
|
|
|
|
result = base.clone()
|
|
target_region = result[tuple(paste_slices)]
|
|
|
|
if opacity_is_tensor:
|
|
cropped_mask = mask[tuple(crop_slices)]
|
|
result[tuple(paste_slices)] = cropped_overlay * cropped_mask + target_region * (1.0 - cropped_mask)
|
|
else:
|
|
if opacity_f >= 1.0:
|
|
result[tuple(paste_slices)] = cropped_overlay
|
|
else:
|
|
result[tuple(paste_slices)] = cropped_overlay * opacity_f + target_region * (1.0 - opacity_f)
|
|
|
|
return result.contiguous()
|
|
|
|
def visitReplaceFunc(self, ctx):
|
|
val = yield ctx.expr(0)
|
|
old = yield ctx.expr(1)
|
|
new = yield ctx.expr(2)
|
|
|
|
if isinstance(val, str):
|
|
return val.replace(str(old), str(new))
|
|
|
|
if self._is_list(val):
|
|
return [new if x == old else x for x in val]
|
|
|
|
if self._is_tensor(val):
|
|
old_t = self._promote_to_tensor(old)
|
|
new_t = self._promote_to_tensor(new)
|
|
return torch.where(val == old_t, new_t, val).contiguous()
|
|
|
|
return val
|
|
|
|
def visitUpperFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if isinstance(val, str):
|
|
return val.upper()
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: upper() requires a string argument")
|
|
|
|
def visitLowerFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
if isinstance(val, str):
|
|
return val.lower()
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: lower() requires a string argument")
|
|
|
|
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):
|
|
return separator.join([str(x) for x in items.flatten().tolist()])
|
|
else:
|
|
return str(items)
|
|
|
|
def visitSubstringFunc(self, ctx):
|
|
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]
|
|
else:
|
|
return string[start_idx:]
|
|
|
|
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):
|
|
val = yield ctx.expr()
|
|
if isinstance(val, str):
|
|
return val.strip()
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: trim() requires a string argument")
|
|
|
|
def visitDilateFunc(self, ctx):
|
|
tsr_val = yield ctx.expr(0)
|
|
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
|
|
tsr = self._promote_to_tensor(tsr_val)
|
|
|
|
original_shape = tsr.shape
|
|
tsr = tsr.float()
|
|
|
|
kernel_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
|
|
|
|
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):
|
|
tsr_val = yield ctx.expr(0)
|
|
tsr = self._promote_to_tensor(tsr_val)
|
|
kernel_size = yield ctx.expr(1) if len(ctx.expr()) > 1 else 3
|
|
tsr = self._promote_to_tensor(tsr_val)
|
|
|
|
original_shape = tsr.shape
|
|
tsr = tsr.float()
|
|
|
|
kernel_size = int(kernel_size.item()) if self._is_tensor(kernel_size) else int(kernel_size)
|
|
|
|
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):
|
|
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):
|
|
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:
|
|
rgb_val = yield ctx.expr(0)
|
|
|
|
use_degrees = False
|
|
if num_args == 2:
|
|
degrees_val = yield ctx.expr(1)
|
|
use_degrees = bool(degrees_val.item() if self._is_tensor(degrees_val) else degrees_val)
|
|
|
|
if self._is_list(rgb_val):
|
|
if len(rgb_val) != 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: rgb_to_hsv expects 3 values [r, g, b], got {len(rgb_val)}")
|
|
r = self._promote_to_tensor(rgb_val[0])
|
|
g = self._promote_to_tensor(rgb_val[1])
|
|
b = self._promote_to_tensor(rgb_val[2])
|
|
else:
|
|
rgb = self._promote_to_tensor(rgb_val)
|
|
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]
|
|
else:
|
|
# Separate r, g, b mode
|
|
r = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
g = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
|
|
|
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)
|
|
|
|
# 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 torch.stack([h, s, v], dim=-1)
|
|
|
|
def visitHsvToRgbFunc(self, ctx):
|
|
num_args = len(ctx.expr())
|
|
|
|
# Determine mode
|
|
if num_args == 1 or num_args == 2:
|
|
hsv_val = yield ctx.expr(0)
|
|
|
|
use_degrees = False
|
|
if num_args == 2:
|
|
degrees_val = yield ctx.expr(1)
|
|
use_degrees = bool(degrees_val.item() if self._is_tensor(degrees_val) else degrees_val)
|
|
|
|
if self._is_list(hsv_val):
|
|
if len(hsv_val) != 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: hsv_to_rgb expects 3 values [h, s, v], got {len(hsv_val)}")
|
|
h = self._promote_to_tensor(hsv_val[0])
|
|
s = self._promote_to_tensor(hsv_val[1])
|
|
v = self._promote_to_tensor(hsv_val[2])
|
|
else:
|
|
hsv = self._promote_to_tensor(hsv_val)
|
|
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]
|
|
else:
|
|
# Separate h, s, v mode
|
|
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
|
|
|
|
# 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 torch.stack([r, g, b], dim=-1)
|
|
|
|
def visitEntropyFunc(self, ctx):
|
|
val = self._promote_to_tensor((yield ctx.expr()))
|
|
# Shannon entropy: -sum(p * log(p))
|
|
p = F.softmax(val.flatten().float(), dim=0)
|
|
entropy = -torch.sum(p * torch.log(p + 1e-10))
|
|
return entropy.item()
|
|
|
|
def visitCorrFunc(self, ctx):
|
|
x = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
y = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
# Pearson correlation coefficient
|
|
vx = x - torch.mean(x)
|
|
vy = y - torch.mean(y)
|
|
corr = torch.sum(vx * vy) / (torch.sqrt(torch.sum(vx ** 2)) * torch.sqrt(torch.sum(vy ** 2)))
|
|
return corr.item()
|
|
|
|
def visitConcatFunc(self, ctx):
|
|
exprs = ctx.expr()
|
|
items = []
|
|
for i in range(len(exprs) - 1):
|
|
items.append((yield exprs[i]))
|
|
|
|
dim_val = yield exprs[-1]
|
|
if any(isinstance(x, str) for x in items):
|
|
return "".join(str(items))
|
|
|
|
if all(self._is_list(x) for x in items):
|
|
res = []
|
|
for x in items:
|
|
res.extend(list(x))
|
|
return res
|
|
|
|
tensors = [self._promote_to_tensor(x) for x in items]
|
|
d = int(dim_val.item()) if self._is_tensor(dim_val) else int(dim_val)
|
|
return torch.cat(tensors, dim=d)
|
|
|
|
def visitIntFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
def as_int(ctx,val):
|
|
if self._is_tensor(val):
|
|
return val.to(torch.int32).contiguous()
|
|
|
|
if self._is_list(val):
|
|
return [as_int(ctx,x) for x in val]
|
|
|
|
if isinstance(val, str):
|
|
return int(float(val))
|
|
|
|
if val is None:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Cannot convert None to a number")
|
|
|
|
return int(float(val))
|
|
return as_int(ctx,val)
|
|
|
|
def visitFloatFunc(self, ctx):
|
|
val = yield ctx.expr()
|
|
def as_float(ctx,val):
|
|
if self._is_tensor(val):
|
|
return val.to(torch.float).contiguous()
|
|
|
|
if self._is_list(val):
|
|
return [as_float(ctx,x) for x in val]
|
|
|
|
if isinstance(val, str):
|
|
return float(val)
|
|
|
|
if val is None:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Cannot convert None to a number")
|
|
|
|
return float(val)
|
|
return as_float(ctx,val)
|
|
|
|
|
|
def i2rgb(self,val):
|
|
if self._is_list(val):
|
|
t = []
|
|
for v in val:
|
|
t.append(self.i2rgb(v))
|
|
return t
|
|
r = (val >> 16) & 0xFF
|
|
g = (val >> 8) & 0xFF
|
|
b = val & 0xFF
|
|
if self._is_tensor(val): return torch.stack([r.to(torch.float)/256, g.to(torch.float)/256, b.to(torch.float)/256], dim=-1).contiguous()
|
|
return [r/256, g/256, b/256]
|
|
|
|
def visitInt_to_rgb(self,ctx):
|
|
val = (yield ctx.expr())
|
|
return self.i2rgb(val)
|
|
|
|
def visitRgb_to_int(self,ctx):
|
|
def clamp255_scaled(x):
|
|
return max(0, min(255, int(float(x) * 256)))
|
|
|
|
if len(ctx.expr()) == 1:
|
|
rgb_val = yield ctx.expr(0)
|
|
if self._is_tensor(rgb_val):
|
|
if rgb_val.shape[-1] != 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: rgb_to_int expects tensor with last dim=3, got shape {rgb_val.shape}")
|
|
r = (rgb_val[..., 0] * 256).clamp(0, 255).to(torch.int32)
|
|
g = (rgb_val[..., 1] * 256).clamp(0, 255).to(torch.int32)
|
|
b = (rgb_val[..., 2] * 256).clamp(0, 255).to(torch.int32)
|
|
return ((r << 16) | (g << 8) | b).contiguous()
|
|
if self._is_list(rgb_val):
|
|
if len(rgb_val) != 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: rgb_to_int expects 3 values [r, g, b], got {len(rgb_val)}")
|
|
return (clamp255_scaled(rgb_val[0]) << 16) | (clamp255_scaled(rgb_val[1]) << 8) | clamp255_scaled(rgb_val[2])
|
|
|
|
r = yield ctx.expr(0)
|
|
g = yield ctx.expr(1)
|
|
b = yield ctx.expr(2)
|
|
|
|
if self._is_tensor(r) or self._is_tensor(g) or self._is_tensor(b):
|
|
r = (self._promote_to_tensor(r) * 256).clamp(0, 255).to(torch.int32)
|
|
g = (self._promote_to_tensor(g) * 256).clamp(0, 255).to(torch.int32)
|
|
b = (self._promote_to_tensor(b) * 256).clamp(0, 255).to(torch.int32)
|
|
return ((r << 16) | (g << 8) | b).contiguous()
|
|
|
|
return (clamp255_scaled(r) << 16) | (clamp255_scaled(g) << 8) | clamp255_scaled(b)
|
|
|
|
def visitLinspaceFunc(self, ctx):
|
|
"""linspace(start, end, steps) - linearly spaced values"""
|
|
start_val = yield ctx.expr(0)
|
|
end_val = yield ctx.expr(1)
|
|
steps_val = yield ctx.expr(2)
|
|
|
|
start = float(start_val.item()) if self._is_tensor(start_val) else float(start_val)
|
|
end = float(end_val.item()) if self._is_tensor(end_val) else float(end_val)
|
|
steps = int(steps_val.item()) if self._is_tensor(steps_val) else int(steps_val)
|
|
|
|
return torch.linspace(start, end, steps, device=self.device)
|
|
|
|
def visitLogspaceFunc(self, ctx):
|
|
"""logspace(start, end, steps, base) - logarithmically spaced values"""
|
|
start_val = yield ctx.expr(0)
|
|
end_val = yield ctx.expr(1)
|
|
steps_val = yield ctx.expr(2)
|
|
base_val = yield ctx.expr(3)
|
|
|
|
start = float(start_val.item()) if self._is_tensor(start_val) else float(start_val)
|
|
end = float(end_val.item()) if self._is_tensor(end_val) else float(end_val)
|
|
base = float(base_val.item()) if self._is_tensor(base_val) else float(base_val)
|
|
steps = int(steps_val.item()) if self._is_tensor(steps_val) else int(steps_val)
|
|
|
|
return torch.logspace(start, end, steps, base=base, device=self.device)
|
|
|
|
def visitRollFunc(self, ctx):
|
|
"""roll(x, shift, [dim]) - circular shift of elements"""
|
|
x = self._promote_to_tensor((yield ctx.expr(0)))
|
|
shift_val = yield ctx.expr(1)
|
|
shift = int(shift_val.item()) if self._is_tensor(shift_val) else int(shift_val)
|
|
|
|
dim = 0
|
|
if len(ctx.expr()) > 2:
|
|
dim_val = yield ctx.expr(2)
|
|
dim = int(dim_val.item()) if self._is_tensor(dim_val) else int(dim_val)
|
|
|
|
return torch.roll(x, shifts=shift, dims=dim)
|
|
|
|
def visitErfFunc(self, ctx):
|
|
"""erf(x) - error function"""
|
|
x = self._promote_to_tensor((yield ctx.expr()))
|
|
return torch.erf(x)
|
|
|
|
def visitErfinvFunc(self, ctx):
|
|
"""erfinv(x) - inverse error function"""
|
|
x = self._promote_to_tensor((yield ctx.expr()))
|
|
return torch.erfinv(x)
|
|
|
|
def visitWhereFunc(self, ctx):
|
|
cond = (yield ctx.expr(0))
|
|
a = (yield ctx.expr(1))
|
|
b = (yield ctx.expr(2))
|
|
|
|
if self._is_tensor(cond) or self._is_tensor(a) or self._is_tensor(b):
|
|
cond_t = self._promote_to_tensor(cond)
|
|
if cond_t.dtype != torch.bool:
|
|
cond_t = cond_t != 0
|
|
a_t = self._promote_to_tensor(a)
|
|
b_t = self._promote_to_tensor(b)
|
|
return torch.where(cond_t, a_t, b_t).contiguous()
|
|
|
|
def rec(c, av, bv):
|
|
if self._is_list(c):
|
|
out = []
|
|
for i, ci in enumerate(c):
|
|
ai = av[i] if self._is_list(av) and i < len(av) else av
|
|
bi = bv[i] if self._is_list(bv) and i < len(bv) else bv
|
|
out.append(rec(ci, ai, bi))
|
|
return out
|
|
return av if bool(c) else bv
|
|
|
|
return rec(cond, a, b)
|
|
|
|
def visitHistogramFunc(self, ctx):
|
|
x = self._promote_to_tensor((yield ctx.expr(0))).float().flatten()
|
|
bins_raw = (yield ctx.expr(1))
|
|
min_raw = (yield ctx.expr(2))
|
|
max_raw = (yield ctx.expr(3))
|
|
|
|
bins = int(bins_raw.item()) if self._is_tensor(bins_raw) else int(bins_raw)
|
|
min_v = float(min_raw.item()) if self._is_tensor(min_raw) else float(min_raw)
|
|
max_v = float(max_raw.item()) if self._is_tensor(max_raw) else float(max_raw)
|
|
|
|
if bins <= 0:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: histogram bins must be > 0")
|
|
if max_v <= min_v:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: histogram requires max > min")
|
|
|
|
return torch.histc(x, bins=bins, min=min_v, max=max_v).contiguous()
|
|
|
|
def visitFlowMagFunc(self, ctx):
|
|
flow = self._promote_to_tensor((yield ctx.expr()))
|
|
if flow.shape[-1] != 2:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: flow_mag expects [..., 2], got {tuple(flow.shape)}")
|
|
dx = flow[..., 0]
|
|
dy = flow[..., 1]
|
|
return torch.sqrt(dx * dx + dy * dy).contiguous()
|
|
|
|
def visitFlowAngFunc(self, ctx):
|
|
flow = self._promote_to_tensor((yield ctx.expr()))
|
|
if flow.shape[-1] != 2:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: flow_ang expects [..., 2], got {tuple(flow.shape)}")
|
|
dx = flow[..., 0]
|
|
dy = flow[..., 1]
|
|
|
|
return torch.atan2(dy, dx).contiguous()
|
|
|
|
|
|
|
|
def _interpolate_impl(self, ctx, mode):
|
|
x = self._promote_to_tensor((yield ctx.expr(0)))
|
|
arg = (yield ctx.expr(1))
|
|
|
|
if x.ndim not in (3, 4, 5):
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: interpolate expects 3D, 4D or 5D input, got shape {tuple(x.shape)}"
|
|
)
|
|
|
|
spatial_dims = x.ndim - 2
|
|
|
|
if mode == "linear":
|
|
interp_mode = {1: "linear", 2: "bilinear", 3: "trilinear"}[spatial_dims]
|
|
else:
|
|
interp_mode = mode
|
|
|
|
if self._is_list(arg):
|
|
size = tuple(self._to_int(v, ctx, "interpolate") for v in arg)
|
|
if len(size) != spatial_dims:
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: interpolate size length {len(size)} does not match spatial dims {spatial_dims}"
|
|
)
|
|
kwargs = {"size": size, "mode": interp_mode}
|
|
|
|
elif self._is_tensor(arg):
|
|
if arg.numel() == 1:
|
|
size = (self._to_int(arg, ctx, "interpolate"),)
|
|
else:
|
|
size = tuple(self._to_int(v, ctx, "interpolate") for v in arg.flatten().tolist())
|
|
|
|
if len(size) != spatial_dims:
|
|
raise ValueError(
|
|
f"{ctx.start.line}:{ctx.start.column}: interpolate size length {len(size)} does not match spatial dims {spatial_dims}"
|
|
)
|
|
kwargs = {"size": size, "mode": interp_mode}
|
|
|
|
else:
|
|
kwargs = {"scale_factor": float(arg), "mode": interp_mode}
|
|
|
|
if interp_mode in ("linear", "bilinear", "trilinear"):
|
|
kwargs["align_corners"] = False
|
|
|
|
return F.interpolate(x, **kwargs)
|
|
|
|
def visitInterpolateLinearFunc(self, ctx):
|
|
return (yield from self._interpolate_impl(ctx, "linear"))
|
|
|
|
def visitInterpolateAreaFunc(self, ctx):
|
|
return (yield from self._interpolate_impl(ctx, "area"))
|
|
|
|
def visitInterpolateNearestExactFunc(self, ctx):
|
|
return (yield from self._interpolate_impl(ctx, "nearest-exact"))
|
|
|
|
def _rgb_triplet_from_arg(self, ctx, val, func_name):
|
|
if self._is_list(val):
|
|
if len(val) != 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {func_name} expects 3 values [r, g, b], got {len(val)}")
|
|
return (
|
|
self._promote_to_tensor(val[0]).float(),
|
|
self._promote_to_tensor(val[1]).float(),
|
|
self._promote_to_tensor(val[2]).float(),
|
|
)
|
|
|
|
rgb = self._promote_to_tensor(val)
|
|
if rgb.shape[-1] != 3:
|
|
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: {func_name} expects tensor with last dim=3, got shape {rgb.shape}")
|
|
rgb = rgb.float()
|
|
return rgb[..., 0], rgb[..., 1], rgb[..., 2]
|
|
|
|
def _srgb_to_linear(self, c):
|
|
return torch.where(c <= 0.04045, c / 12.92, torch.pow((c + 0.055) / 1.055, 2.4))
|
|
|
|
def _linear_to_oklab(self, r, g, b):
|
|
rgb = torch.stack([r, g, b], dim=-1).float()
|
|
rgb = self._srgb_to_linear(rgb)
|
|
|
|
l = 0.4122214708 * rgb[..., 0] + 0.5363325363 * rgb[..., 1] + 0.0514459929 * rgb[..., 2]
|
|
m = 0.2119034982 * rgb[..., 0] + 0.6806995451 * rgb[..., 1] + 0.1073969566 * rgb[..., 2]
|
|
s = 0.0883024619 * rgb[..., 0] + 0.2817188376 * rgb[..., 1] + 0.6299787005 * rgb[..., 2]
|
|
|
|
l_ = torch.pow(l, 1.0 / 3.0)
|
|
m_ = torch.pow(m, 1.0 / 3.0)
|
|
s_ = torch.pow(s, 1.0 / 3.0)
|
|
|
|
L = 0.2104542553 * l_ + 0.7936177850 * m_ - 0.0040720468 * s_
|
|
a = 1.9779984951 * l_ - 2.4285922050 * m_ + 0.4505937099 * s_
|
|
b = 0.0259040371 * l_ + 0.7827717662 * m_ - 0.8086757660 * s_
|
|
return torch.stack([L, a, b], dim=-1)
|
|
|
|
def _linear_to_cielab(self, r, g, b):
|
|
rgb = torch.stack([r, g, b], dim=-1).float()
|
|
rgb = self._srgb_to_linear(rgb)
|
|
|
|
x = 0.4124564 * rgb[..., 0] + 0.3575761 * rgb[..., 1] + 0.1804375 * rgb[..., 2]
|
|
y = 0.2126729 * rgb[..., 0] + 0.7151522 * rgb[..., 1] + 0.0721750 * rgb[..., 2]
|
|
z = 0.0193339 * rgb[..., 0] + 0.1191920 * rgb[..., 1] + 0.9503041 * rgb[..., 2]
|
|
|
|
xn, yn, zn = 0.95047, 1.0, 1.08883
|
|
delta = 6.0 / 29.0
|
|
|
|
def f(t):
|
|
return torch.where(t > delta ** 3, torch.pow(t, 1.0 / 3.0), t / (3.0 * delta ** 2) + 4.0 / 29.0)
|
|
|
|
fx = f(x / xn)
|
|
fy = f(y / yn)
|
|
fz = f(z / zn)
|
|
|
|
L = 116.0 * fy - 16.0
|
|
a = 500.0 * (fx - fy)
|
|
b = 200.0 * (fy - fz)
|
|
return torch.stack([L, a, b], dim=-1)
|
|
|
|
def visitRgbToOklabFunc(self, ctx):
|
|
num_args = len(ctx.expr())
|
|
|
|
if num_args == 1:
|
|
rgb_val = yield ctx.expr(0)
|
|
r, g, b = self._rgb_triplet_from_arg(ctx, rgb_val, "rgb_to_oklab")
|
|
else:
|
|
r = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
g = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
|
|
|
return self._linear_to_oklab(r, g, b)
|
|
|
|
def visitRgbToCielabFunc(self, ctx):
|
|
num_args = len(ctx.expr())
|
|
|
|
if num_args == 1:
|
|
rgb_val = yield ctx.expr(0)
|
|
r, g, b = self._rgb_triplet_from_arg(ctx, rgb_val, "rgb_to_cielab")
|
|
else:
|
|
r = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
g = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
|
return self._linear_to_cielab(r, g, b)/100
|
|
|
|
def _linear_to_srgb(self, c):
|
|
return torch.where(c <= 0.0031308, 12.92 * c, 1.055 * torch.pow(c, 1.0 / 2.4) - 0.055)
|
|
|
|
def _oklab_to_rgb(self, L, a, b):
|
|
l_ = L + 0.3963377774 * a + 0.2158037573 * b
|
|
m_ = L - 0.1055613458 * a - 0.0638541728 * b
|
|
s_ = L - 0.0894841775 * a - 1.2914855480 * b
|
|
|
|
l = l_ * l_ * l_
|
|
m = m_ * m_ * m_
|
|
s = s_ * s_ * s_
|
|
|
|
r_lin = +4.0767416621 * l - 3.3077115913 * m + 0.2309699292 * s
|
|
g_lin = -1.2684380046 * l + 2.6097574011 * m - 0.3413193965 * s
|
|
b_lin = -0.0041960863 * l - 0.7034186147 * m + 1.7076147010 * s
|
|
|
|
rgb_lin = torch.stack([r_lin, g_lin, b_lin], dim=-1)
|
|
return self._linear_to_srgb(rgb_lin)
|
|
|
|
def _cielab_to_rgb(self, L, a, b):
|
|
delta = 6.0 / 29.0
|
|
|
|
fy = (L + 16.0) / 116.0
|
|
fx = fy + a / 500.0
|
|
fz = fy - b / 200.0
|
|
|
|
def finv(t):
|
|
return torch.where(t > delta, t * t * t, 3.0 * (delta ** 2) * (t - 4.0 / 29.0))
|
|
|
|
xn, yn, zn = 0.95047, 1.0, 1.08883
|
|
x = xn * finv(fx)
|
|
y = yn * finv(fy)
|
|
z = zn * finv(fz)
|
|
|
|
r_lin = 3.2404542 * x - 1.5371385 * y - 0.4985314 * z
|
|
g_lin = -0.9692660 * x + 1.8760108 * y + 0.0415560 * z
|
|
b_lin = 0.0556434 * x - 0.2040259 * y + 1.0572252 * z
|
|
|
|
rgb_lin = torch.stack([r_lin, g_lin, b_lin], dim=-1)
|
|
return self._linear_to_srgb(rgb_lin)
|
|
|
|
def visitOklabToRgbFunc(self, ctx):
|
|
num_args = len(ctx.expr())
|
|
if num_args == 1:
|
|
lab_val = yield ctx.expr(0)
|
|
L, a, b = self._rgb_triplet_from_arg(ctx, lab_val, "oklab_to_rgb")
|
|
else:
|
|
L = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
a = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
|
|
|
return self._oklab_to_rgb(L, a, b)
|
|
|
|
def visitCielabToRgbFunc(self, ctx):
|
|
num_args = len(ctx.expr())
|
|
if num_args == 1:
|
|
lab_val = yield ctx.expr(0)
|
|
L, a, b = self._rgb_triplet_from_arg(ctx, lab_val, "cielab_to_rgb")
|
|
else:
|
|
L = self._promote_to_tensor((yield ctx.expr(0))).float()
|
|
a = self._promote_to_tensor((yield ctx.expr(1))).float()
|
|
b = self._promote_to_tensor((yield ctx.expr(2))).float()
|
|
|
|
return self._cielab_to_rgb(L*100, a*100, b*100)
|
|
|
|
def visitLambdaExp(self, ctx):
|
|
# ctx.paramList() nebo None; body je ctx.block() nebo ctx.expr()
|
|
params = []
|
|
if ctx.paramList():
|
|
params = [node.getText() for node in ctx.paramList().VARIABLE()]
|
|
body = ctx.block() if ctx.block() else ctx.expr()
|
|
closure = self.variables.copy()
|
|
return LambdaFunction(params, body, closure)
|
|
|
|
def visitRidgedFunc(self, ctx):
|
|
"""ridged(seed, scale, [octaves], [offset], [shape])"""
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
|
|
scale_val = yield ctx.expr(1)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
|
|
octaves = 4
|
|
expr_idx = 2
|
|
if len(ctx.expr()) > expr_idx:
|
|
oct_val = yield ctx.expr(expr_idx)
|
|
octaves = int(oct_val.item()) if self._is_tensor(oct_val) else int(oct_val)
|
|
expr_idx += 1
|
|
|
|
offset = None
|
|
if len(ctx.expr()) > expr_idx:
|
|
offset_val = yield ctx.expr(expr_idx)
|
|
offset = offset_val
|
|
expr_idx += 1
|
|
|
|
shape = self.shape
|
|
if len(ctx.expr()) > expr_idx:
|
|
shape_arg = (yield ctx.expr(expr_idx))
|
|
if self._is_tensor(shape_arg):
|
|
shape = tuple(shape_arg.long().flatten().tolist())
|
|
elif self._is_list(shape_arg):
|
|
shape = tuple(int(x) for x in shape_arg)
|
|
else:
|
|
shape = (int(shape_arg),)
|
|
|
|
if len(shape) == 0:
|
|
return torch.tensor(0.0, device=self.device)
|
|
|
|
offset_list = None
|
|
if offset is not None:
|
|
if self._is_tensor(offset):
|
|
offset_list = [float(x) for x in offset.flatten().tolist()]
|
|
elif self._is_list(offset):
|
|
offset_list = [float(x) for x in offset]
|
|
else:
|
|
offset_list = [float(offset)]
|
|
|
|
# Build 1D ranges with offsets and avoid full meshgrid when some dims == 1
|
|
ranges = [
|
|
torch.arange(s, dtype=torch.float32, device=self.device)
|
|
+ (offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0)
|
|
for i, s in enumerate(shape)
|
|
]
|
|
valid_indices = [i for i, s in enumerate(shape) if s > 1]
|
|
if len(valid_indices) == 0:
|
|
# all dims are size 1: create cheap singleton tensors per-dim
|
|
grids = tuple(
|
|
torch.tensor(
|
|
[offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0],
|
|
dtype=torch.float32,
|
|
device=self.device,
|
|
)
|
|
for i in range(len(shape))
|
|
)
|
|
grids_optimized = grids
|
|
else:
|
|
# only build meshgrid for dimensions with size > 1
|
|
ranges_filtered = [ranges[i] for i in valid_indices]
|
|
grids_filtered = torch.meshgrid(*ranges_filtered, indexing='ij')
|
|
grids_optimized = tuple(grids_filtered)
|
|
|
|
noise = NoiseUtils.ridged_noise_nd(grids_optimized, scale, seed, self.device, octaves=octaves)
|
|
target_shape = tuple(int(x) for x in shape)
|
|
return noise.view(target_shape)
|
|
|
|
def visitDomainWarpFunc(self, ctx):
|
|
"""domain_warp(seed, scale, warp_scale, warp_strength, [octaves], [warp_octaves], [offset], [shape])"""
|
|
seed_val = yield ctx.expr(0)
|
|
seed = int(seed_val.item()) if self._is_tensor(seed_val) else int(seed_val)
|
|
|
|
scale_val = yield ctx.expr(1)
|
|
scale = float(scale_val.item()) if self._is_tensor(scale_val) else float(scale_val)
|
|
|
|
warp_scale_val = yield ctx.expr(2)
|
|
warp_scale = float(warp_scale_val.item()) if self._is_tensor(warp_scale_val) else float(warp_scale_val)
|
|
|
|
warp_strength_val = yield ctx.expr(3)
|
|
warp_strength = float(warp_strength_val.item()) if self._is_tensor(warp_strength_val) else float(warp_strength_val)
|
|
|
|
octaves = 4
|
|
warp_octaves = 2
|
|
expr_idx = 4
|
|
|
|
if len(ctx.expr()) > expr_idx:
|
|
oct_val = yield ctx.expr(expr_idx)
|
|
octaves = int(oct_val.item()) if self._is_tensor(oct_val) else int(oct_val)
|
|
expr_idx += 1
|
|
|
|
if len(ctx.expr()) > expr_idx:
|
|
warp_oct_val = yield ctx.expr(expr_idx)
|
|
warp_octaves = int(warp_oct_val.item()) if self._is_tensor(warp_oct_val) else int(warp_oct_val)
|
|
expr_idx += 1
|
|
|
|
offset = None
|
|
if len(ctx.expr()) > expr_idx:
|
|
offset_val = yield ctx.expr(expr_idx)
|
|
offset = offset_val
|
|
expr_idx += 1
|
|
|
|
shape = self.shape
|
|
if len(ctx.expr()) > expr_idx:
|
|
shape_arg = (yield ctx.expr(expr_idx))
|
|
if self._is_tensor(shape_arg):
|
|
shape = tuple(shape_arg.long().flatten().tolist())
|
|
elif self._is_list(shape_arg):
|
|
shape = tuple(int(x) for x in shape_arg)
|
|
else:
|
|
shape = (int(shape_arg),)
|
|
|
|
if len(shape) == 0:
|
|
return torch.tensor(0.0, device=self.device)
|
|
|
|
offset_list = None
|
|
if offset is not None:
|
|
if self._is_tensor(offset):
|
|
offset_list = [float(x) for x in offset.flatten().tolist()]
|
|
elif self._is_list(offset):
|
|
offset_list = [float(x) for x in offset]
|
|
else:
|
|
offset_list = [float(offset)]
|
|
|
|
# Build 1D ranges with offsets and avoid full meshgrid when some dims == 1
|
|
ranges = [
|
|
torch.arange(s, dtype=torch.float32, device=self.device)
|
|
+ (offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0)
|
|
for i, s in enumerate(shape)
|
|
]
|
|
valid_indices = [i for i, s in enumerate(shape) if s > 1]
|
|
if len(valid_indices) == 0:
|
|
# all dims are size 1: create cheap singleton tensors per-dim
|
|
grids = tuple(
|
|
torch.tensor(
|
|
[offset_list[i] if offset_list is not None and i < len(offset_list) else 0.0],
|
|
dtype=torch.float32,
|
|
device=self.device,
|
|
)
|
|
for i in range(len(shape))
|
|
)
|
|
grids_optimized = grids
|
|
else:
|
|
# only build meshgrid for dimensions with size > 1
|
|
ranges_filtered = [ranges[i] for i in valid_indices]
|
|
grids_filtered = torch.meshgrid(*ranges_filtered, indexing='ij')
|
|
grids_optimized = tuple(grids_filtered)
|
|
|
|
noise = NoiseUtils.domain_warp_noise_nd(
|
|
grids_optimized,
|
|
scale=scale,
|
|
seed=seed,
|
|
device=self.device,
|
|
warp_scale=warp_scale,
|
|
warp_strength=warp_strength,
|
|
octaves=octaves,
|
|
warp_octaves=warp_octaves,
|
|
)
|
|
target_shape = tuple(int(x) for x in shape)
|
|
return noise.view(target_shape)
|