Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ddf36c68c5 | ||
|
|
50d23bdfc4 | ||
|
|
d30e31833f | ||
|
|
d8646f2a36 | ||
|
|
92fd99a357 | ||
|
|
939fb805a1 | ||
|
|
5bd85aaf76 | ||
|
|
36454f9160 | ||
|
|
464cffdfea | ||
|
|
2589b2f129 | ||
|
|
cc66b62c08 | ||
|
|
d0724ccec3 | ||
|
|
1ad2f2d8b2 |
@@ -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
|
||||
|
||||
@@ -32,6 +32,7 @@ if (
|
||||
".comfyui_string_list",
|
||||
".comfyui_uncrop",
|
||||
".comfyui_rc",
|
||||
".comfyui_sound",
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
+284
-2
@@ -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)
|
||||
|
||||
@@ -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}}
|
||||
Reference in New Issue
Block a user