feat: ✨ add TensorOps

pretty rough for now, inspired by blender math nodes
This commit is contained in:
Mel Massadian
2025-02-14 23:47:48 +01:00
parent edcb3da08b
commit a8cf4650ff
+182
View File
@@ -677,6 +677,187 @@ class MTB_ConcatImages:
return (concatenated,)
class MTB_TensorOps:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"tensor": ("IMAGE",),
"operation": (
[
"multiply",
"divide",
"add",
"subtract",
"power",
"clamp",
"abs",
"log",
"exp",
"convert_dtype",
"normalize_range",
"normalize_per_channel",
],
{"default": "multiply"},
),
"value": (
"FLOAT",
{
"default": 1.0,
"min": -1000000.0,
"max": 1000000.0,
"step": 0.01,
},
),
"source_min": (
"FLOAT",
{
"default": 0.0,
"min": -1000000.0,
"max": 1000000.0,
"step": 0.01,
},
),
"source_max": (
"FLOAT",
{
"default": 1.0,
"min": -1000000.0,
"max": 1000000.0,
"step": 0.01,
},
),
"target_min": (
"FLOAT",
{
"default": 0.0,
"min": -1000000.0,
"max": 1000000.0,
"step": 0.01,
},
),
"target_max": (
"FLOAT",
{
"default": 16.0,
"min": -1000000.0,
"max": 1000000.0,
"step": 0.01,
},
),
"dtype": (
["uint8", "float32", "float16", "bfloat16"],
{"default": "float32"},
),
"use_mean": ("BOOLEAN", {"default": False}),
},
"optional": {
"target_tensor": ("IMAGE",),
},
}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "apply"
CATEGORY = "tensor_ops"
def apply(
self,
tensor,
operation="multiply",
value=1.0,
source_min=0.0,
source_max=1.0,
target_min=0.0,
target_max=1.0,
dtype="float32",
use_mean=False,
target_tensor=None,
):
log.debug(
f"Input tensor stats: shape={tensor.shape}, dtype={tensor.dtype}, range=[{tensor.min().item():.6f}, {tensor.max().item():.6f}]"
)
if operation == "normalize_per_channel":
if target_tensor is None:
raise ValueError(
"Target tensor required for per-channel normalization"
)
result = tensor.clone()
for c in range(tensor.shape[-1]):
if use_mean:
source_mean = tensor[..., c].mean()
target_mean = target_tensor[..., c].mean()
scale = target_mean / source_mean
result[..., c] = tensor[..., c] * scale
else:
source_min = tensor[..., c].min()
source_max = tensor[..., c].max()
target_min = target_tensor[..., c].min()
target_max = target_tensor[..., c].max()
normalized = (tensor[..., c] - source_min) / (
source_max - source_min
)
result[..., c] = (
normalized * (target_max - target_min) + target_min
)
log.debug(
f"Channel {c} - Scale: source=[{source_min:.6f}, {source_max:.6f}], target=[{target_min:.6f}, {target_max:.6f}]"
)
elif operation == "normalize_range":
if target_tensor is not None:
target_min = target_tensor.min().item()
target_max = target_tensor.max().item()
log.debug(
f"Using target tensor range: [{target_min:.6f}, {target_max:.6f}]"
)
normalized = (tensor - source_min) / (source_max - source_min)
result = normalized * (target_max - target_min) + target_min
elif operation == "convert_dtype":
if dtype == "float32":
result = tensor.float()
elif dtype == "float16":
result = tensor.half()
elif dtype == "bfloat16":
result = tensor.bfloat16()
else:
result = tensor
if operation == "multiply":
result = tensor * value
elif operation == "divide":
result = tensor / value if value != 0 else tensor
elif operation == "add":
result = tensor + value
elif operation == "subtract":
result = tensor - value
elif operation == "power":
result = torch.pow(tensor, value)
elif operation == "clamp":
if target_tensor is not None:
result = torch.clamp(
tensor,
target_tensor.min().item(),
target_tensor.max().item(),
)
else:
result = torch.clamp(tensor, source_min, source_max)
elif operation == "abs":
result = torch.abs(tensor)
elif operation == "log":
result = torch.log(tensor.clamp(min=1e-10))
elif operation == "exp":
result = torch.exp(tensor)
log.debug(
f"Output tensor stats: shape={result.shape}, dtype={result.dtype}, range=[{result.min().item():.6f}, {result.max().item():.6f}]"
)
return (result,)
__nodes__ = [
MTB_StringReplace,
MTB_FitNumber,
@@ -691,4 +872,5 @@ __nodes__ = [
MTB_FloatsToFloat,
MTB_FloatToFloats,
MTB_FloatsToInts,
MTB_TensorOps,
]