diff --git a/README.md b/README.md index 2878b13..3e0dfbc 100644 --- a/README.md +++ b/README.md @@ -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 diff --git a/more_math/Parser/UnifiedMathVisitor.py b/more_math/Parser/UnifiedMathVisitor.py index dca9ccf..89ce7c8 100644 --- a/more_math/Parser/UnifiedMathVisitor.py +++ b/more_math/Parser/UnifiedMathVisitor.py @@ -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 diff --git a/web/script_text_input.js b/web/script_text_input.js index 91f998f..fcbfce9 100644 --- a/web/script_text_input.js +++ b/web/script_text_input.js @@ -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 = {