13 Commits
4 changed files with 640 additions and 2 deletions
+3
View File
@@ -52,6 +52,9 @@ comfyui_image_ops
JWImageResizeByFactor: Image Resize by Factor
JWImageResizeByShorterSide: Image Resize by Shorter Side
JWImageResizeByLongerSide: Image Resize by Longer Side
JWImageResizeToClosestSDXLResolution: Image Resize to Closest SDXL Resolution
JWImageLoadRGBFromClipboard: Image Load RGB From Clipboard
JWImageLoadRGBA From Clipboard: Image Load RGBA From Clipboard
comfyui_primitive_ops
JWInteger: Integer
+1
View File
@@ -32,6 +32,7 @@ if (
".comfyui_string_list",
".comfyui_uncrop",
".comfyui_rc",
".comfyui_sound",
]
)
+284 -2
View File
@@ -1,11 +1,13 @@
import json
import math
import os
from pathlib import Path
from typing import Optional, Union
import numpy as np
import torch
import torchvision.transforms.functional as F
from PIL import Image
from PIL import Image, ImageGrab
from PIL.PngImagePlugin import PngInfo
from torchvision.transforms import InterpolationMode
@@ -298,7 +300,7 @@ class _:
INPUT_TYPES = lambda: {
"required": {
"image": ("IMAGE",),
"direction": (("horizontal", "vertical"), {"default": "hotizontal"}),
"direction": (("horizontal", "vertical"), {"default": "horizontal"}),
}
}
RETURN_TYPES = ("IMAGE",)
@@ -564,3 +566,283 @@ class _:
image = image.permute(0, 2, 3, 1)
return (image,)
@register_node(
"JWImageResizeToClosestSDXLResolution", "Image Resize to Closest SDXL Resolution"
)
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {
"required": {
"image": ("IMAGE",),
"interpolation_mode": (
["bicubic", "bilinear", "nearest", "nearest exact"],
),
}
}
RETURN_TYPES = ("IMAGE", "INT", "INT")
RETURN_NAMES = ("IMAGE", "WIDTH", "HEIGHT")
FUNCTION = "execute"
# tuples of (height x width)
SDXL_RESOLUTIONS = (
(1024, 1024),
(1152, 896),
(896, 1152),
(1216, 832),
(832, 1216),
(1344, 768),
(768, 1344),
(1536, 640),
(640, 1536),
)
@staticmethod
def compare_fn(img_w: int, img_h: int, resolution: tuple[int, int]):
img_deg = math.atan(img_h / img_w)
xl_deg = math.atan(resolution[0] / resolution[1])
return abs(img_deg - xl_deg)
def execute(
self,
image: torch.Tensor,
interpolation_mode: str,
):
interpolation_mode = interpolation_mode.upper().replace(" ", "_")
interpolation_mode = getattr(InterpolationMode, interpolation_mode)
_, h, w, _ = image.shape
closest_resolution = min(
self.SDXL_RESOLUTIONS, key=lambda res: self.compare_fn(w, h, res)
)
image = image.permute(0, 3, 1, 2)
image = F.resize(
image,
closest_resolution, # type: ignore
interpolation=interpolation_mode, # type: ignore
antialias=True,
)
image = image.permute(0, 2, 3, 1)
return (image, closest_resolution[1], closest_resolution[0])
@register_node(
"JWImageCropToClosestSDXLResolution", "Image Crop to Closest SDXL Resolution"
)
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {
"required": {
"image": ("IMAGE",),
"interpolation_mode": (
["bicubic", "bilinear", "nearest", "nearest exact"],
),
}
}
RETURN_TYPES = ("IMAGE", "INT", "INT")
RETURN_NAMES = ("IMAGE", "WIDTH", "HEIGHT")
FUNCTION = "execute"
# tuples of (height x width)
SDXL_RESOLUTIONS = (
(1024, 1024),
(1152, 896),
(896, 1152),
(1216, 832),
(832, 1216),
(1344, 768),
(768, 1344),
(1536, 640),
(640, 1536),
)
@staticmethod
def angle(w: int, h: int):
return math.atan(h / w)
@staticmethod
def compare_fn(img_w: int, img_h: int, resolution: tuple[int, int]):
img_deg = math.atan(img_h / img_w)
xl_deg = math.atan(resolution[0] / resolution[1])
return abs(img_deg - xl_deg)
def execute(
self,
image: torch.Tensor,
interpolation_mode: str,
):
interpolation_mode = interpolation_mode.upper().replace(" ", "_")
interpolation_mode = getattr(InterpolationMode, interpolation_mode)
_, h, w, _ = image.shape
closest_resolution = min(
self.SDXL_RESOLUTIONS, key=lambda res: self.compare_fn(w, h, res)
)
img_deg = self.angle(w, h)
target_deg = self.angle(closest_resolution[1], closest_resolution[0])
if img_deg > target_deg:
# image is taller and narrower than target
w_scaled = closest_resolution[1]
h_scaled = max(round(closest_resolution[1] / w * h), 0)
else:
# image is wider and shorter than target
h_scaled = closest_resolution[0]
w_scaled = max(round(closest_resolution[0] / h * w), 0)
image = image.permute(0, 3, 1, 2)
image = F.resize(
image,
[h_scaled, w_scaled],
interpolation=interpolation_mode, # type: ignore
antialias=True,
)
image = F.center_crop(
image,
closest_resolution, # type: ignore
)
image = image.permute(0, 2, 3, 1)
return (image, closest_resolution[1], closest_resolution[0])
@register_node("JWImageResizeToMegapixels", "Image Resize to Megapixels")
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {
"required": {
"image": ("IMAGE",),
"megapixels": (
"FLOAT",
{"default": 1.0, "min": 0.01, "step": 0.01, "max": 99999.0},
),
"divisible_by": ("INT", {"default": 32, "min": 1, "step": 1, "max": 99999}),
"interpolation_mode": (
["bicubic", "bilinear", "nearest", "nearest exact"],
),
}
}
RETURN_TYPES = ("IMAGE", "INT", "INT")
RETURN_NAMES = ("IMAGE", "WIDTH", "HEIGHT")
FUNCTION = "execute"
def execute(
self,
image: torch.Tensor,
megapixels: float,
divisible_by: int,
interpolation_mode: str,
):
assert isinstance(image, torch.Tensor)
assert isinstance(megapixels, float)
assert isinstance(divisible_by, int)
assert isinstance(interpolation_mode, str)
interpolation_mode = interpolation_mode.upper().replace(" ", "_")
interpolation_mode = getattr(InterpolationMode, interpolation_mode)
_, h, w, _ = image.shape
# find target resolution without rounding
aspect_ratio = w / h
target_h_exact = math.sqrt(megapixels * 1_000_000 / aspect_ratio)
target_w_exact = math.sqrt(megapixels * 1_000_000 * aspect_ratio)
# round to nearest `divisible_by` (ensure it doesn't round to 0)
final_h = max(round(target_h_exact / divisible_by) * divisible_by, divisible_by)
final_w = max(round(target_w_exact / divisible_by) * divisible_by, divisible_by)
# find scale amount needed to fully cover final dimensions
scale_w = final_w / w
scale_h = final_h / h
scale = max(scale_w, scale_h)
# calculate resize dimensions (prevent floating point rounding into smaller dimension)
resize_w = max(round(w * scale), final_w)
resize_h = max(round(h * scale), final_h)
image = image.permute(0, 3, 1, 2)
image = F.resize(
image,
[resize_h, resize_w],
interpolation=interpolation_mode, # type: ignore
antialias=True,
)
image = F.center_crop(
image,
[final_h, final_w], # type: ignore
)
image = image.permute(0, 2, 3, 1)
return (image, final_w, final_h)
def get_image_from_clipboard(rgba=False) -> Optional[torch.Tensor]:
rv = ImageGrab.grabclipboard()
if rv is None:
return None
if isinstance(rv, list):
if len(rv) == 0:
return None
img = Image.open(rv[0]).convert("RGBA" if rgba else "RGB")
else:
# rv is some kind of image
img = rv.convert("RGBA" if rgba else "RGB")
img = np.array(img).astype(np.float32) / 255.0
img = torch.from_numpy(img).unsqueeze(0)
return img
@register_node("JWImageLoadRGBFromClipboard", "Image Load RGB From Clipboard")
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {"required": {}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "execute"
def execute(self):
img = get_image_from_clipboard(rgba=False)
if img is None:
raise ValueError(f"failed to get image from clipboard")
return (img,)
def IS_CHANGED(self, *args):
# This value will be compared with previous 'IS_CHANGED' outputs
# If inequal, then this node will be considered as modified
return get_image_from_clipboard(rgba=False)
@register_node("JWImageLoadRGBA From Clipboard", "Image Load RGBA From Clipboard")
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {"required": {}}
RETURN_TYPES = ("IMAGE", "MASK")
FUNCTION = "execute"
def execute(self):
img = get_image_from_clipboard(rgba=True)
if img is None:
raise ValueError(f"failed to get image from clipboard")
color = img[:, :, :, 0:3]
mask = img[0, :, :, 3]
mask = 1 - mask # invert mask
return (color, mask)
def IS_CHANGED(self, *args):
# This value will be compared with previous 'IS_CHANGED' outputs
# If inequal, then this node will be considered as modified
return get_image_from_clipboard(rgba=True)
+352
View File
@@ -0,0 +1,352 @@
import json
import math
import os
from pathlib import Path
from typing import Literal, TypedDict
import soundfile as sf
import torch
import torchaudio
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
def register_node(identifier: str, display_name: str):
def decorator(cls):
NODE_CLASS_MAPPINGS[identifier] = cls
NODE_DISPLAY_NAME_MAPPINGS[identifier] = display_name
return cls
return decorator
FLOAT_MAX = 99999999999999999.0
class Audio(TypedDict):
sample_rate: int
# .shape => torch.Size([1, 2, 1403182])
waveform: torch.Tensor
def scalar_to_db(scalar: float):
return 20 * math.log10(scalar)
def db_to_scalar(db: float):
return 10 ** (db / 20)
def load_audio(
path: str | Path,
sr: None | int | float = None,
offset: float = 0.0,
duration: float | None = None,
make_stereo: bool = True,
) -> Audio:
import librosa
mix, sr = librosa.load(path, sr=sr, mono=False, offset=offset, duration=duration)
mix = torch.from_numpy(mix)
# If stereo, shape will be:
# torch.Size([2, 1403182])
# If mono, shape will be:
# torch.Size([1403182])
#
# Ensure shape is [channels, data]
if len(mix.shape) == 1:
mix = torch.stack([mix], dim=0)
assert len(mix.shape) == 2
# Convert mono to stereo if needed
if make_stereo:
if mix.shape[0] == 1:
mix = torch.cat([mix, mix], dim=0)
elif mix.shape[0] == 2:
pass
else:
raise ValueError(
f"Input audio has {mix.shape[0]} channels, cannot convert to stereo (2 channels)"
)
# Add extra dimension for batch size
# shape => torch.Size([2, 1403182])
mix = torch.unsqueeze(mix, 0)
# shape => torch.Size([1, 2, 1403182])
return {
"sample_rate": round(sr),
"waveform": mix,
}
def save_audio(path: str | Path, mix: torch.Tensor, sr):
path = str(path)
# make sure tensor has shape [channels, data]
if len(mix.shape) == 3:
if mix.shape[0] > 1:
raise ValueError("Audio batch size is more than 1")
mix = mix[0]
elif len(mix.shape) == 2:
pass
elif len(mix.shape) == 1:
mix = torch.unsqueeze(mix, 0)
else:
raise ValueError(f"Invalid tensor shape: {mix.shape}")
subtype = "FLOAT" if path.lower().endswith("wav") else None
sf.write(path, mix.T, sr, subtype=subtype)
def write_audio_comment(path: str | Path, comment: str):
try:
from mediafile import MediaFile
except ImportError as e:
print(
"[WARN] Failed to import `mediafile`, saved audio files will not have metadata"
)
return
f = MediaFile(path)
f.comments = comment
f.save()
@register_node("JWLoadAudio", "Audio Load")
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {
"required": {
"path": ("STRING", {"default": "./audio.mp3"}),
"gain_db": ("FLOAT", {"default": 0, "min": -100, "max": 100}),
"offset_seconds": ("FLOAT", {"default": 0, "min": 0, "max": FLOAT_MAX}),
"duration_seconds": ("FLOAT", {"default": 0, "min": 0, "max": FLOAT_MAX}),
"resample_to_hz": ("FLOAT", {"default": 0, "min": 0, "max": FLOAT_MAX}),
"make_stereo": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("AUDIO",)
FUNCTION = "execute"
def execute(
self,
path: str,
gain_db: float,
offset_seconds: float,
duration_seconds: float,
resample_to_hz: float,
make_stereo: bool,
) -> tuple[Audio]:
rv = load_audio(
path,
sr=resample_to_hz if resample_to_hz > 0 else None,
offset=offset_seconds,
duration=duration_seconds if duration_seconds > 0 else None,
make_stereo=make_stereo,
)
if gain_db != 0.0:
gain_scalar = db_to_scalar(gain_db)
rv["waveform"] = gain_scalar * rv["waveform"]
return (rv,)
@classmethod
def IS_CHANGED(
cls,
path: str,
*args,
):
if os.path.exists(path):
mtime = os.path.getmtime(path)
else:
mtime = None
return (mtime, path, *args)
@register_node("JWAudioBlend", "Audio Blend")
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {
"required": {
"a": ("AUDIO",),
"b": ("AUDIO",),
"ratio": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0}),
"if_durations_differ": (
("use_longest", "use_shortest"),
{"default": "use_longest"},
),
"if_samplerates_differ": (
("use_highest", "use_lowest"),
{"default": "use_highest"},
),
}
}
RETURN_TYPES = ("AUDIO",)
FUNCTION = "execute"
def execute(
self,
a: Audio,
b: Audio,
ratio: float,
if_durations_differ: Literal["use_longest", "use_shortest"],
if_samplerates_differ: Literal["use_highest", "use_lowest"],
) -> tuple[Audio]:
import librosa
# shallow clone audios
a = {**a}
b = {**b}
# # if they have different batch size, attempt to resolve them
# if a["waveform"].shape[0] != b["waveform"].shape[0]:
# pass
# if they have different channels, attempt to resolve them
if a["waveform"].shape[1] != b["waveform"].shape[1]:
# if one of them is mono, distribute it
if a["waveform"].shape[1] == 1:
a["waveform"] = a["waveform"].expand(-1, b["waveform"].shape[1])
elif b["waveform"].shape[1] == 1:
b["waveform"] = b["waveform"].expand(-1, a["waveform"].shape[1])
# ensure audio has same sample rate
sr = a["sample_rate"]
if a["sample_rate"] != b["sample_rate"]:
# determine which rate to use
if if_samplerates_differ == "use_highest":
sr = max(a["sample_rate"], b["sample_rate"])
elif if_samplerates_differ == "use_lowest":
sr = min(a["sample_rate"], b["sample_rate"])
else:
raise NotImplementedError(if_samplerates_differ)
# do the resampling
if a["sample_rate"] != sr:
a["waveform"] = torchaudio.functional.resample(
a["waveform"], a["sample_rate"], sr
)
if b["sample_rate"] != sr:
b["waveform"] = torchaudio.functional.resample(
b["waveform"], b["sample_rate"], sr
)
# ensure input has same length
duration = a["waveform"].shape[-1]
if a["waveform"].shape[-1] != b["waveform"].shape[-1]:
# determine which duration to use
if if_durations_differ == "use_longest":
duration = max(a["waveform"].shape[-1], b["waveform"].shape[-1])
elif if_durations_differ == "use_shortest":
duration = min(a["waveform"].shape[-1], b["waveform"].shape[-1])
else:
raise NotImplementedError(if_samplerates_differ)
def waveform_with_duration(wave: torch.Tensor, new_duration: int):
batch, channels, original_duration = wave.shape
if original_duration >= new_duration:
return wave[:, :, :new_duration]
else:
rv = torch.zeros(batch, channels, new_duration)
rv[:, :, :original_duration] = wave[:, :, :original_duration]
return rv
# do the chopping
if a["waveform"].shape[-1] != duration:
a["waveform"] = waveform_with_duration(a["waveform"], duration)
if b["waveform"].shape[-1] != duration:
b["waveform"] = waveform_with_duration(b["waveform"], duration)
rv: Audio = {
"sample_rate": sr,
"waveform": a["waveform"] * (1.0 - ratio) + b["waveform"] * ratio,
}
return (rv,)
class ResultItem(TypedDict):
filename: str
subfolder: str
type: Literal["output"]
@register_node("JWAudioSaveToPath", "Audio Save to Path")
class _:
CATEGORY = "jamesWalker55"
INPUT_TYPES = lambda: {
"required": {
"audio": ("AUDIO",),
"path": ("STRING", {"default": "./audio.mp3"}),
"overwrite": ("BOOLEAN", {"default": True}),
},
"hidden": {
"prompt": "PROMPT",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
RETURN_TYPES = ()
FUNCTION = "execute"
OUTPUT_NODE = True
def execute(
self,
path: str | Path,
audio: Audio,
overwrite: bool,
prompt=None,
extra_pnginfo=None,
):
path = Path(path)
path.parent.mkdir(exist_ok=True)
metadata = {**(extra_pnginfo or {})}
if prompt is not None:
metadata["prompt"] = prompt
metadata_str = json.dumps(metadata)
results: list[ResultItem] = []
if audio["waveform"].shape[0] == 1:
# batch has 1 audio only
if overwrite or not path.exists():
save_audio(
path,
audio["waveform"][0],
audio["sample_rate"],
)
write_audio_comment(path, metadata_str)
results.append(
{
"filename": path.name,
"subfolder": str(path.parent),
"type": "output",
}
)
else:
# batch has multiple images
for i, subwaveform in enumerate(audio["waveform"]):
subpath = path.with_stem(f"{path.stem}-{i}")
if overwrite or not path.exists():
save_audio(
subpath,
subwaveform,
audio["sample_rate"],
)
write_audio_comment(subpath, metadata_str)
results.append(
{
"filename": subpath.name,
"subfolder": str(subpath.parent),
"type": "output",
}
)
return {"ui": {"audio": results}}