AI: add string functions

This commit is contained in:
mcDandy
2026-02-26 12:46:13 +01:00
parent ed030f1542
commit d80c8c1156
3 changed files with 90 additions and 763 deletions
+10 -49
View File
@@ -116,7 +116,7 @@ You can also get the node from comfy manager under the name of More math.
### Aggregates & Tensor Operations
- `tmin(x, y)`: Element-wise minimum of x and y.
- `tmax(x, y): Element-wise maximum of x and y.
- `tmax(x, y)`: Element-wise maximum of x and y.
- `smin(x, ...)`: Scalar minimum. Returns the single smallest value across all input tensors/values.
- `smax(x, ...)`: Scalar maximum. Returns the single largest value across all input tensors/values.
- `sum(x)`: Sum of all elements.
@@ -207,58 +207,19 @@ You can also get the node from comfy manager under the name of More math.
- `nan_to_num(x, nan_value, posinf_value, neginf_value)` or `nvl`: Replaces NaN and infinite values in tensor with specified values.
- `remap(v, i_min, i_max, o_min, o_max)`: Remaps value `v` from input range `[i_min, i_max]` to output range `[o_min, o_max]`.
- `timestamp()` or `now`: Returns current UNIX timestamp (precision to microseconds, can be different on other systems)
- `count(x)` or `length(x)` or `cnt(x)`: Returns the length of a list or the size of the first dimension of a tensor.
- `count(x)` or `length(x)` or `cnt(x)`: Returns the length of a list, string, or the size of the first dimension of a tensor.
### Random Distributions
Generates random noise with default shape of aither first input or maximum of input sizes, depending on node setting.
### String Operations
- `random_normal(seed,[shape])` or `randn` or `noise`: generates a random tensor with normal distribution (var=1, mean=0).
- `random_uniform(seed,[shape])` or `rand`: generates a random tensor with uniform distribution [0, 1).
- `random_exponential(seed, lambda,[shape])` or `rande`: generates a random tensor with exponential distribution.
- `random_cauchy(seed, median, sigma,[shape])` or `randc`: generates a random tensor with Cauchy distribution.
- `random_log_normal(seed, mean, std,[shape])` or `randln`: generates a random tensor with log-normal distribution.
- `random_bernoulli(seed, p,[shape])` or `randb`: generates a random tensor with Bernoulli distribution. Parameter `p` is the probability of getting 1, can be aither float or tensor. If p is tensor, shape is ignored.
- `random_poisson(seed, lambda,[shape])` or `randp`: generates a random tensor with Poisson distribution. Lambda can be either float or tensor.
- `random_gamma(seed, shape, scale,[shape])` or `randg`: generates a random tensor with Gamma distribution. Shape parameter (α) controls the shape, scale parameter (θ) controls the scale.
- `random_beta(seed, alpha, beta,[shape])` or `randbeta`: generates a random tensor with Beta distribution in range [0, 1]. Alpha and beta are shape parameters.
- `random_laplace(seed, loc, scale,[shape])` or `randl`: generates a random tensor with Laplace (double exponential) distribution. Useful for L1 regularization and robust statistics.
- `random_gumbel(seed, loc, scale,[shape])` or `randgumbel`: generates a random tensor with Gumbel distribution. Used in Gumbel-softmax trick for neural networks.
- `random_weibull(seed, scale, concentration,[shape])` or `randw`: generates a random tensor with Weibull distribution. Used in reliability analysis and survival modeling.
- `random_chi2(seed, df,[shape])` or `randchi2`: generates a random tensor with Chi-squared distribution. Degrees of freedom `df` controls the shape. Sum of squared normal distributions.
- `random_studentt(seed, df,[shape])` or `randt`: generates a random tensor with Student's t distribution. Has heavier tails than normal distribution, useful for robust noise. As `df` increases, approaches normal distribution.
### Noise Generation
- `perlin(seed, scale, [octaves,[offset, [shape]]])` or `perlin_noise`: generates Perlin noise. `scale` controls the frequency of the noise, `octaves` adds additional layers of noise, `offset` offsets the noise pattern, `shape` controls the output shape (default is determined by node inputs and settings).
- `plasma(seed, scale, [octaves,[offset, [shape]]])` or `turbulence` or `plasma_noise`: generates Plasma noise. Same parameters as perlin noise.
- `voronoi(seed, scale, [jitter], [offset], [shape])` or `voronoi_noise`: generates Voronoi noise. `scale` controls the frequency of the noise, `jitter` adds randomness to the cell boundaries, `offset` offsets the noise pattern, `shape` controls the output shape (default is determined by node inputs and settings).
### Bitwise Operations
Bitwise operations work with scalars, tensors, and lists, preserving bit patterns (especially important for floats where bit patterns are preserved, not values converted).
#### Shift Operators
- `a << b`: Left shift operator. Shifts bits of `a` left by `b` positions.
- `a >> b`: Right shift operator. Shifts bits of `a` right by `b` positions.
#### Bitwise Functions
- `band(a, b)` or `bitwise_and(a, b)`: Bitwise AND. Returns bits set in both operands.
- `bor(a, b)` or `bitwise_or(a, b)`: Bitwise OR. Returns bits set in either operand.
- `bxor(a, b)` or `bitwise_xor(a, b)`: Bitwise XOR. Returns bits set in exactly one operand.
- `bnot(a)` or `bitwise_not(a)`: Bitwise NOT. Inverts all bits in the operand.
- `bitcount(a)`, `popcount(a)`, or `popcnt(a)`: Count set bits. Returns the number of set bits (1s) in the binary representation as a float.
### Stack
- `stack_push(id, value)`: Pushes value to stack with id.
- `stack_pop(id)`: Pops value from stack with id.
- `stack_get(id)`: Gets value from stack with id.
- `stack_clear(id)`: Clears stack with id.
- `stack_has(id)`: Checks if stack with id exists.
- `upper(str)`: Converts string to uppercase.
- `lower(str)`: Converts string to lowercase.
- `split(str, [delimiter])`: Splits string into a list. Default delimiter is space.
- `join(list, [separator])`: Joins list elements into a string. Default separator is empty string.
- `substring(str, start, [length])` or `substr`: Extracts substring starting at `start` position. If length is omitted, extracts to end.
- `find(str, search)`: Returns position of first occurrence of `search` in `str`, or -1 if not found.
- `trim(str)`: Removes leading and trailing whitespace from string.
## Variables
+78 -713
View File
@@ -1110,6 +1110,8 @@ class UnifiedMathVisitor(MathExprVisitor):
def visitCountFunc(self, ctx):
val = yield ctx.expr()
if isinstance(val, str):
return float(len(val))
if self._is_list(val):
return float(len(val))
if self._is_tensor(val):
@@ -1784,7 +1786,6 @@ class UnifiedMathVisitor(MathExprVisitor):
return assigned_val
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
@@ -2225,719 +2226,83 @@ class UnifiedMathVisitor(MathExprVisitor):
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 visitRgbToHsvFunc(self, ctx):
r = self._promote_to_tensor((yield ctx.expr(0)))
g = self._promote_to_tensor((yield ctx.expr(1)))
b = self._promote_to_tensor((yield ctx.expr(2)))
# Ensure values are in [0, 1]
r = torch.clamp(r, 0, 1)
g = torch.clamp(g, 0, 1)
b = torch.clamp(b, 0, 1)
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 calculation
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
# Saturation
s = torch.where(max_rgb > 0, diff / max_rgb, torch.zeros_like(max_rgb))
# Value
v = max_rgb
return [h, s, v]
def visitHsvToRgbFunc(self, ctx):
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)))
h = h % 360
s = torch.clamp(s, 0, 1)
v = torch.clamp(v, 0, 1)
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]
return [r + m, g + m, b + m]
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_tensor(shape_val):
shape = shape_val.int().tolist()
elif self._is_list(shape_val):
shape = [int(float(x)) for x in shape_val]
else:
shape = [int(float(shape_val))]
return torch.full(shape, value, device=self.device,dtype=type)
def visitSoftmaxFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
return F.softmax(val.float())
def visitSoftminFunc(self, ctx):
val = self._promote_to_tensor((yield ctx.expr()))
return F.softmax(-val.float())
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:
# Cross product requires vectors with 3-component last dimension (supports broadcasting)
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")
return torch.cross(a, b, dim=-1)
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."""
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())
# 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'
)
noise = NoiseUtils.perlin_noise_nd(grids, scale, seed, self.device)
if octaves > 1:
result = noise
amplitude = 0.5
frequency = 2.0
for oct in range(octaves - 1):
scaled_grids = tuple(g * frequency for g in grids)
octave_noise = NoiseUtils.perlin_noise_nd(scaled_grids, scale / frequency, seed + oct, self.device)
result = result + octave_noise * amplitude
amplitude *= 0.5
frequency *= 2.0
noise = result / (2 - 2**(-octaves))
return noise
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 = 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'
)
# Call perlin_noise_nd with all coordinate grids
noise = NoiseUtils.plasma_noise_nd(grids, scale, seed, self.device)
# Apply octaves (fBm-like composition)
if octaves > 1:
result = noise
amplitude = 0.5
frequency = 2.0
for oct in range(octaves - 1):
scaled_grids = tuple(g * frequency for g in grids)
octave_noise = NoiseUtils.plasma_noise_nd(scaled_grids, scale / frequency, seed + oct, self.device)
result = result + octave_noise * amplitude
amplitude *= 0.5
frequency *= 2.0
noise = result / (2 - 2**(-octaves))
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 visitOverlayFunc(self, ctx):
base = yield ctx.expr(0)
overlay = yield ctx.expr(1)
offset_raw = yield ctx.expr(2)
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:]
offset = 0
end = min(len(base), offset + len(overlay))
overlay_len = end - offset
return base[:offset] + overlay[:overlay_len] + base[end:]
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:]
offset = 0
result = list(base)
end = min(len(base), offset + len(overlay))
for i, val in enumerate(overlay[:end - offset]):
result[offset + i] = val
return result
# Handle tensors (existing implementation)
base = self._promote_to_tensor(base)
overlay = self._promote_to_tensor(overlay)
offset = offset_raw
# Convert offset to list of ints
if self._is_tensor(offset):
offset = [int(x) for x in offset.flatten().tolist()]
elif self._is_list(offset):
offset = [int(x) for x in offset]
else:
offset = [int(offset)]
# Ensure offset matches base dimensions
if len(offset) != base.ndim:
raise ValueError(f"{ctx.start.line}:{ctx.start.column}: Offset dimensions {len(offset)} must match base dimensions {base.ndim}")
# Calculate crop and paste regions
crop_slices = []
paste_slices = []
for i in range(base.ndim):
off = offset[i]
overlay_size = overlay.shape[i]
base_size = base.shape[i]
if off >= base_size:
return base # Overlay outside of base, return original
# Determine overlay crop region (what part of overlay to use)
crop_start = max(0, -off) # Crop from overlay if offset is negative
crop_end = min(overlay_size, base_size - off) # Crop if overlay extends beyond base
# Determine base paste region (where to place overlay in base)
paste_start = max(0, off) # Start position in base
paste_end = min(base_size, off + overlay_size) # End position in base
crop_slices.append(slice(crop_start, crop_end))
paste_slices.append(slice(paste_start, paste_end))
# Crop overlay to fit
cropped_overlay = overlay[tuple(crop_slices)]
# Create result by cloning base and pasting overlay
result = base.clone()
result[tuple(paste_slices)] = cropped_overlay
return result
+2 -1
View File
@@ -24,7 +24,8 @@ const FUNCTIONS = new Set([
"random_exponential", "randln", "random_log_normal", "randb", "random_bernoulli", "randp", "random_poisson", "randg",
"random_gamma", "randbeta", "random_beta", "randl", "random_laplace", "randgumbel", "random_gumbel", "randw",
"random_weibull", "randchi2", "random_chi2", "randt", "random_studentt", "perlin", "perlin_noise", "cellular", "voronoi",
"worley", "cellular_noise", "voronoi_noise", "plasma", "turbulence", "plasma_noise"
"worley", "cellular_noise", "voronoi_noise", "plasma", "turbulence", "plasma_noise",
"upper", "lower", "split", "join", "substring", "substr", "find", "trim", "replace"
]);
const BRACKET_PAIRS = {