Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
210dc072b1 | ||
|
|
9bd02bd62b | ||
|
|
a445c82b2b | ||
|
|
91eceb7d57 | ||
|
|
090418eb8d | ||
|
|
f1830ba85b | ||
|
|
51ed4a0bc4 | ||
|
|
c62b79b6d9 | ||
|
|
1138a1f4d9 | ||
|
|
f37981bffb | ||
|
|
e40e9244b3 | ||
|
|
c01aafb508 | ||
|
|
61171a2a87 | ||
|
|
91b4761689 | ||
|
|
75b47d8eb4 | ||
|
|
c6b58caca6 | ||
|
|
8965cfb3c3 | ||
|
|
6df1bc2298 | ||
|
|
6107619468 | ||
|
|
7dbb8fb35c | ||
|
|
746c1cefbd | ||
|
|
e5e027e25c | ||
|
|
dd6673b4d5 | ||
|
|
42c116dbb0 | ||
|
|
12ee9ebe16 | ||
|
|
12aa820fdc | ||
|
|
beae36b7ed | ||
|
|
b886928282 | ||
|
|
1721ff7a70 | ||
|
|
8e83109a4e | ||
|
|
e4510faffb | ||
|
|
27c0905ba4 | ||
|
|
896a59a294 | ||
|
|
b3c5a98603 | ||
|
|
52d1b5c874 | ||
|
|
e3b2493b1f | ||
|
|
20952fd8c1 | ||
|
|
43461905b2 | ||
|
|
a2aaa32fbc | ||
|
|
707517937a | ||
|
|
71722e4c7f | ||
|
|
0d7bcc5e23 | ||
|
|
2503bc2cf3 | ||
|
|
3a6a0c2b52 | ||
|
|
efaa2dcd6c | ||
|
|
c77a9b386e | ||
|
|
c3bacdc0c4 | ||
|
|
4e97ff8c4a | ||
|
|
64fa05980d | ||
|
|
d78b709e31 | ||
|
|
4d6caa301b | ||
|
|
ad57177bcd | ||
|
|
fc00f4a094 | ||
|
|
633352bf5d | ||
|
|
3e97c544f2 | ||
|
|
3bf0cfa0cc | ||
|
|
50abaace75 | ||
|
|
8d538c9678 |
@@ -7,15 +7,19 @@ on:
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'sipherxyz' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@main
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
MIT License
|
||||
|
||||
Copyright (c) 2025 VIXION
|
||||
|
||||
Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
of this software and associated documentation files (the "Software"), to deal
|
||||
in the Software without restriction, including without limitation the rights
|
||||
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
copies of the Software, and to permit persons to whom the Software is
|
||||
furnished to do so, subject to the following conditions:
|
||||
|
||||
The above copyright notice and this permission notice shall be included in all
|
||||
copies or substantial portions of the Software.
|
||||
|
||||
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
||||
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
||||
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
||||
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
||||
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
||||
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
|
||||
SOFTWARE.
|
||||
@@ -129,89 +129,27 @@ Remove objects from images using LaMa model.
|
||||
|
||||
### LLM Nodes
|
||||
|
||||
#### LLMApiConfig
|
||||
- **LLM API Config**: Generic model settings (model, tokens, temperature).
|
||||
- **NanoBanana API Config**: Preset for Gemini 2.5 Flash Image (modalities/aspect ratio).
|
||||
- **OpenAI API**: Connect to OpenAI chat/completions.
|
||||
- **OpenRouter API**: Route to many providers via OpenRouter; supports text and images where available.
|
||||
- **Gemini API**: Google Gemini (text and image generation).
|
||||
- **Claude API**: Anthropic Claude messages.
|
||||
- **AWS Bedrock Claude API**: Claude via AWS Bedrock.
|
||||
- **AWS Bedrock Mistral API**: Mistral completions via AWS Bedrock.
|
||||
- **LLM Message**: Build message lists (system/user/assistant) with optional images.
|
||||
- **LLM Chat**: Run multi-turn chats; returns text and optional images.
|
||||
- **LLM Completion**: Single-prompt completion.
|
||||
|
||||
Configures generic LLM API parameters.
|
||||

|
||||
|
||||
**Inputs:**
|
||||

|
||||
|
||||
- `model`: Model name (GPT-3.5, GPT-4, etc)
|
||||
- `max_token`: Maximum tokens
|
||||
- `temperature`: Temperature parameter
|
||||
# Known Issues
|
||||
|
||||
#### OpenAIApi
|
||||
## AV_controlnetPreprocessor is missing
|
||||
|
||||
Configures OpenAI API access.
|
||||
`AV_controlnetPreprocessor` is a wrapper for [comfyui_controlnet_aux](https://github.com/Fannovel16/comfyui_controlnet_aux) that I created to quickly switch between multiple preprocessors, instead of having a separate node for each one. It requires `comfyui_controlnet_aux` to be installed, otherwise it will not be available.
|
||||
|
||||
**Inputs:**
|
||||
Since `AIO_Preprocessor` is already implemented in `comfyui_controlnet_aux`, this node will be deprecated. You are recommended to switch to using `AIO_Preprocessor` node directly.
|
||||
|
||||
- `openai_api_key`: OpenAI API key
|
||||
- `endpoint`: API endpoint URL
|
||||
|
||||
### Claude API Nodes
|
||||
|
||||
#### ClaudeApi
|
||||
|
||||
Configures Anthropic Claude API access.
|
||||
|
||||
**Inputs:**
|
||||
|
||||
- `claude_api_key`: Claude API key
|
||||
- `endpoint`: API endpoint
|
||||
- `version`: API version
|
||||
|
||||
#### AwsBedrockClaudeApi
|
||||
|
||||
Configures AWS Bedrock Claude API access.
|
||||
|
||||
**Inputs:**
|
||||
|
||||
- `aws_access_key_id`: AWS access key
|
||||
- `aws_secret_access_key`: AWS secret key
|
||||
- `region`: AWS region
|
||||
- `version`: API version
|
||||
|
||||
#### AwsBedrockMistralApi
|
||||
|
||||
Configures AWS Bedrock Mistral API access.
|
||||
|
||||
**Inputs:**
|
||||
|
||||
- `aws_access_key_id`: AWS access key
|
||||
- `aws_secret_access_key`: AWS secret key
|
||||
- `region`: AWS region
|
||||
|
||||
#### LLMMessage
|
||||
|
||||
Creates a message for LLM interaction.
|
||||
|
||||
**Inputs:**
|
||||
|
||||
- `role`: Message role (system/user/assistant)
|
||||
- `text`: Message content
|
||||
- `image`: Optional image input
|
||||
- `messages`: Previous message history
|
||||
|
||||
#### LLMChat
|
||||
|
||||
Handles chat interactions with LLMs.
|
||||
|
||||
**Inputs:**
|
||||
|
||||
- `messages`: Message history
|
||||
- `api`: LLM API configuration
|
||||
- `config`: Model configuration
|
||||
- `seed`: Random seed
|
||||
|
||||
#### LLMCompletion
|
||||
|
||||
Handles completion requests to LLMs.
|
||||
|
||||
**Inputs:**
|
||||
|
||||
- `prompt`: Input prompt
|
||||
- `api`: LLM API configuration
|
||||
- `config`: Model configuration
|
||||
- `seed`: Random seed
|
||||
|
||||

|
||||
|
||||
@@ -66,7 +66,7 @@ class AVControlNetLoader(ControlNetLoader):
|
||||
|
||||
RETURN_TYPES = ("CONTROL_NET",)
|
||||
FUNCTION = "load_controlnet"
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def load_controlnet(self, control_net_name, control_net_override="None", timestep_keyframe=None):
|
||||
return load_controlnet(control_net_name, control_net_override, timestep_keyframe=timestep_keyframe)
|
||||
@@ -90,7 +90,8 @@ class AV_ControlNetPreprocessor:
|
||||
RETURN_TYPES = ("IMAGE", "STRING")
|
||||
RETURN_NAMES = ("IMAGE", "CNET_NAME")
|
||||
FUNCTION = "detect_controlnet"
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
DESCRIPTION = "DEPRECATED: Use comfyui_controlnet_aux's AIO Preprocessor instead"
|
||||
|
||||
def detect_controlnet(self, image, preprocessor, sd_version, resolution=512, preprocessor_override="None"):
|
||||
if preprocessor_override != "None":
|
||||
@@ -136,7 +137,7 @@ class AVControlNetEfficientStacker:
|
||||
RETURN_TYPES = ("CONTROL_NET_STACK",)
|
||||
RETURN_NAMES = ("CNET_STACK",)
|
||||
FUNCTION = "control_net_stacker"
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def control_net_stacker(
|
||||
self,
|
||||
@@ -233,7 +234,7 @@ class AVControlNetEfficientLoader(ControlNetApply):
|
||||
|
||||
RETURN_TYPES = ("CONDITIONING",)
|
||||
FUNCTION = "load_controlnet"
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def load_controlnet(
|
||||
self,
|
||||
@@ -289,7 +290,7 @@ class AVControlNetEfficientLoaderAdvanced(ControlNetApplyAdvanced):
|
||||
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||
RETURN_NAMES = ("positive", "negative")
|
||||
FUNCTION = "load_controlnet"
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def load_controlnet(
|
||||
self,
|
||||
@@ -333,5 +334,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AV_ControlNetEfficientLoaderAdvanced": "ControlNet Loader Adv.",
|
||||
"AV_ControlNetEfficientStacker": "ControlNet Stacker Adv.",
|
||||
"AV_ControlNetEfficientStackerSimple": "ControlNet Stacker",
|
||||
"AV_ControlNetPreprocessor": "ControlNet Preprocessor",
|
||||
"AV_ControlNetPreprocessor": "[Deprecated] ControlNet Preprocessor",
|
||||
}
|
||||
|
||||
@@ -8,7 +8,7 @@ import comfy.controlnet
|
||||
from ..utils import load_module
|
||||
|
||||
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
||||
advanced_cnet_dir_names = ["AdvancedControlNet", "ComfyUI-Advanced-ControlNet"]
|
||||
advanced_cnet_dir_names = ["AdvancedControlNet", "ComfyUI-Advanced-ControlNet", "comfyui-advanced-controlnet"]
|
||||
|
||||
|
||||
def comfy_load_controlnet(control_net_name: str, **_):
|
||||
|
||||
@@ -48,7 +48,7 @@ _preprocessors_map = {
|
||||
"teed": "TEEDPreprocessor",
|
||||
"color": "ColorPreprocessor",
|
||||
"sam": "SAMPreprocessor",
|
||||
"tile": "TilePreprocessor"
|
||||
"tile": "TilePreprocessor",
|
||||
}
|
||||
|
||||
|
||||
@@ -67,10 +67,10 @@ try:
|
||||
break
|
||||
|
||||
if module_path is None:
|
||||
raise Exception("Could not find ControlNetPreprocessors nodes")
|
||||
raise Exception("Could not find comfyui_controlnet_aux nodes, AV_ControlNetPreprocessor will not work. Please install comfyui_controlnet_aux first")
|
||||
|
||||
module = load_module(module_path)
|
||||
print("Loaded ControlNetPreprocessors nodes from", module_path)
|
||||
print("Loaded comfyui_controlnet_aux nodes from", module_path)
|
||||
|
||||
nodes: Dict = getattr(module, "NODE_CLASS_MAPPINGS")
|
||||
available_preprocessors: list[str] = getattr(module, "PREPROCESSOR_OPTIONS")
|
||||
|
||||
@@ -20,7 +20,7 @@ class KSamplerWithSharpness(KSampler):
|
||||
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Sampling"
|
||||
CATEGORY = "ArtVenture/Sampling"
|
||||
|
||||
def sample(self, *args, sharpness=2.0, **kwargs):
|
||||
patch.sharpness = sharpness
|
||||
@@ -43,7 +43,7 @@ class KSamplerAdvancedWithSharpness(KSamplerAdvanced):
|
||||
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Sampling"
|
||||
CATEGORY = "ArtVenture/Sampling"
|
||||
|
||||
def sample(self, *args, sharpness=2.0, **kwargs):
|
||||
patch.sharpness = sharpness
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
import torch
|
||||
from typing import Union, Optional
|
||||
|
||||
|
||||
Tensor = torch.Tensor
|
||||
@@ -7,12 +8,12 @@ Dtype = torch.Type
|
||||
pad = torch.nn.functional.pad
|
||||
|
||||
|
||||
def _compute_zero_padding(kernel_size: tuple[int, int] | int) -> tuple[int, int]:
|
||||
def _compute_zero_padding(kernel_size: Union[tuple[int, int], int]) -> tuple[int, int]:
|
||||
ky, kx = _unpack_2d_ks(kernel_size)
|
||||
return (ky - 1) // 2, (kx - 1) // 2
|
||||
|
||||
|
||||
def _unpack_2d_ks(kernel_size: tuple[int, int] | int) -> tuple[int, int]:
|
||||
def _unpack_2d_ks(kernel_size: Union[tuple[int, int], int]) -> tuple[int, int]:
|
||||
if isinstance(kernel_size, int):
|
||||
ky = kx = kernel_size
|
||||
else:
|
||||
@@ -26,17 +27,14 @@ def _unpack_2d_ks(kernel_size: tuple[int, int] | int) -> tuple[int, int]:
|
||||
|
||||
def gaussian(
|
||||
window_size: int,
|
||||
sigma: Tensor | float,
|
||||
sigma: Union[Tensor, float],
|
||||
*,
|
||||
device: Device | None = None,
|
||||
dtype: Dtype | None = None,
|
||||
device: Optional[Device] = None,
|
||||
dtype: Optional[Dtype] = None,
|
||||
) -> Tensor:
|
||||
batch_size = sigma.shape[0]
|
||||
|
||||
x = (
|
||||
torch.arange(window_size, device=sigma.device, dtype=sigma.dtype)
|
||||
- window_size // 2
|
||||
).expand(batch_size, -1)
|
||||
x = (torch.arange(window_size, device=sigma.device, dtype=sigma.dtype) - window_size // 2).expand(batch_size, -1)
|
||||
|
||||
if window_size % 2 == 0:
|
||||
x = x + 0.5
|
||||
@@ -48,68 +46,58 @@ def gaussian(
|
||||
|
||||
def get_gaussian_kernel1d(
|
||||
kernel_size: int,
|
||||
sigma: float | Tensor,
|
||||
sigma: Union[float, Tensor],
|
||||
force_even: bool = False,
|
||||
*,
|
||||
device: Device | None = None,
|
||||
dtype: Dtype | None = None,
|
||||
device: Optional[Device] = None,
|
||||
dtype: Optional[Dtype] = None,
|
||||
) -> Tensor:
|
||||
return gaussian(kernel_size, sigma, device=device, dtype=dtype)
|
||||
|
||||
|
||||
def get_gaussian_kernel2d(
|
||||
kernel_size: tuple[int, int] | int,
|
||||
sigma: tuple[float, float] | Tensor,
|
||||
kernel_size: Union[tuple[int, int], int],
|
||||
sigma: Union[tuple[float, float], Tensor],
|
||||
force_even: bool = False,
|
||||
*,
|
||||
device: Device | None = None,
|
||||
dtype: Dtype | None = None,
|
||||
device: Optional[Device] = None,
|
||||
dtype: Optional[Dtype] = None,
|
||||
) -> Tensor:
|
||||
sigma = torch.Tensor([[sigma, sigma]]).to(device=device, dtype=dtype)
|
||||
|
||||
ksize_y, ksize_x = _unpack_2d_ks(kernel_size)
|
||||
sigma_y, sigma_x = sigma[:, 0, None], sigma[:, 1, None]
|
||||
|
||||
kernel_y = get_gaussian_kernel1d(
|
||||
ksize_y, sigma_y, force_even, device=device, dtype=dtype
|
||||
)[..., None]
|
||||
kernel_x = get_gaussian_kernel1d(
|
||||
ksize_x, sigma_x, force_even, device=device, dtype=dtype
|
||||
)[..., None]
|
||||
kernel_y = get_gaussian_kernel1d(ksize_y, sigma_y, force_even, device=device, dtype=dtype)[..., None]
|
||||
kernel_x = get_gaussian_kernel1d(ksize_x, sigma_x, force_even, device=device, dtype=dtype)[..., None]
|
||||
|
||||
return kernel_y * kernel_x.view(-1, 1, ksize_x)
|
||||
|
||||
|
||||
def _bilateral_blur(
|
||||
input: Tensor,
|
||||
guidance: Tensor | None,
|
||||
kernel_size: tuple[int, int] | int,
|
||||
sigma_color: float | Tensor,
|
||||
sigma_space: tuple[float, float] | Tensor,
|
||||
guidance: Union[Tensor, None],
|
||||
kernel_size: Union[tuple[int, int], int],
|
||||
sigma_color: Union[float, Tensor],
|
||||
sigma_space: Union[tuple[float, float], Tensor],
|
||||
border_type: str = "reflect",
|
||||
color_distance_type: str = "l1",
|
||||
) -> Tensor:
|
||||
if isinstance(sigma_color, Tensor):
|
||||
sigma_color = sigma_color.to(device=input.device, dtype=input.dtype).view(
|
||||
-1, 1, 1, 1, 1
|
||||
)
|
||||
sigma_color = sigma_color.to(device=input.device, dtype=input.dtype).view(-1, 1, 1, 1, 1)
|
||||
|
||||
ky, kx = _unpack_2d_ks(kernel_size)
|
||||
pad_y, pad_x = _compute_zero_padding(kernel_size)
|
||||
|
||||
padded_input = pad(input, (pad_x, pad_x, pad_y, pad_y), mode=border_type)
|
||||
unfolded_input = (
|
||||
padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2)
|
||||
) # (B, C, H, W, Ky x Kx)
|
||||
unfolded_input = padded_input.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
|
||||
|
||||
if guidance is None:
|
||||
guidance = input
|
||||
unfolded_guidance = unfolded_input
|
||||
else:
|
||||
padded_guidance = pad(guidance, (pad_x, pad_x, pad_y, pad_y), mode=border_type)
|
||||
unfolded_guidance = (
|
||||
padded_guidance.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2)
|
||||
) # (B, C, H, W, Ky x Kx)
|
||||
unfolded_guidance = padded_guidance.unfold(2, ky, 1).unfold(3, kx, 1).flatten(-2) # (B, C, H, W, Ky x Kx)
|
||||
|
||||
diff = unfolded_guidance - guidance.unsqueeze(-1)
|
||||
if color_distance_type == "l1":
|
||||
@@ -118,13 +106,9 @@ def _bilateral_blur(
|
||||
color_distance_sq = diff.square().sum(1, keepdim=True)
|
||||
else:
|
||||
raise ValueError("color_distance_type only acceps l1 or l2")
|
||||
color_kernel = (
|
||||
-0.5 / sigma_color**2 * color_distance_sq
|
||||
).exp() # (B, 1, H, W, Ky x Kx)
|
||||
color_kernel = (-0.5 / sigma_color**2 * color_distance_sq).exp() # (B, 1, H, W, Ky x Kx)
|
||||
|
||||
space_kernel = get_gaussian_kernel2d(
|
||||
kernel_size, sigma_space, device=input.device, dtype=input.dtype
|
||||
)
|
||||
space_kernel = get_gaussian_kernel2d(kernel_size, sigma_space, device=input.device, dtype=input.dtype)
|
||||
space_kernel = space_kernel.view(-1, 1, 1, 1, kx * ky)
|
||||
|
||||
kernel = space_kernel * color_kernel
|
||||
@@ -134,9 +118,9 @@ def _bilateral_blur(
|
||||
|
||||
def bilateral_blur(
|
||||
input: Tensor,
|
||||
kernel_size: tuple[int, int] | int = (13, 13),
|
||||
sigma_color: float | Tensor = 3.0,
|
||||
sigma_space: tuple[float, float] | Tensor = 3.0,
|
||||
kernel_size: Union[tuple[int, int], int] = (13, 13),
|
||||
sigma_color: Union[float, Tensor] = 3.0,
|
||||
sigma_space: Union[tuple[float, float], Tensor] = 3.0,
|
||||
border_type: str = "reflect",
|
||||
color_distance_type: str = "l1",
|
||||
) -> Tensor:
|
||||
@@ -154,9 +138,9 @@ def bilateral_blur(
|
||||
def joint_bilateral_blur(
|
||||
input: Tensor,
|
||||
guidance: Tensor,
|
||||
kernel_size: tuple[int, int] | int,
|
||||
sigma_color: float | Tensor,
|
||||
sigma_space: tuple[float, float] | Tensor,
|
||||
kernel_size: Union[tuple[int, int], int],
|
||||
sigma_color: Union[float, Tensor],
|
||||
sigma_space: Union[tuple[float, float], Tensor],
|
||||
border_type: str = "reflect",
|
||||
color_distance_type: str = "l1",
|
||||
) -> Tensor:
|
||||
@@ -174,9 +158,9 @@ def joint_bilateral_blur(
|
||||
class _BilateralBlur(torch.nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
kernel_size: tuple[int, int] | int,
|
||||
sigma_color: float | Tensor,
|
||||
sigma_space: tuple[float, float] | Tensor,
|
||||
kernel_size: Union[tuple[int, int], int],
|
||||
sigma_color: Union[float, Tensor],
|
||||
sigma_space: Union[tuple[float, float], Tensor],
|
||||
border_type: str = "reflect",
|
||||
color_distance_type: str = "l1",
|
||||
) -> None:
|
||||
|
||||
@@ -16,14 +16,10 @@ try:
|
||||
module_path = None
|
||||
|
||||
for custom_node in custom_nodes:
|
||||
custom_node = (
|
||||
custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
|
||||
)
|
||||
custom_node = custom_node if not os.path.islink(custom_node) else os.readlink(custom_node)
|
||||
for module_dir in efficieny_dir_names:
|
||||
if module_dir in os.listdir(custom_node):
|
||||
module_path = os.path.abspath(
|
||||
os.path.join(custom_node, module_dir)
|
||||
)
|
||||
module_path = os.path.abspath(os.path.join(custom_node, module_dir))
|
||||
break
|
||||
|
||||
if module_path is None:
|
||||
@@ -49,7 +45,7 @@ try:
|
||||
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Sampling"
|
||||
CATEGORY = "ArtVenture/Sampling"
|
||||
|
||||
def sample(self, *args, sharpness=2.0, **kwargs):
|
||||
patch.sharpness = sharpness
|
||||
@@ -69,7 +65,7 @@ try:
|
||||
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Sampling"
|
||||
CATEGORY = "ArtVenture/Sampling"
|
||||
|
||||
def sampleadv(self, *args, sharpness=2.0, **kwargs):
|
||||
patch.sharpness = sharpness
|
||||
@@ -87,7 +83,7 @@ try:
|
||||
inputs["optional"]["lora_override"] = ("STRING", {"default": "None"})
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def efficientloader(
|
||||
self,
|
||||
@@ -108,9 +104,7 @@ try:
|
||||
if lora_override != "None":
|
||||
lora_name = lora_override
|
||||
|
||||
return super().efficientloader(
|
||||
ckpt_name, vae_name, clip_skip, lora_name, *args, **kwargs
|
||||
)
|
||||
return super().efficientloader(ckpt_name, vae_name, clip_skip, lora_name, *args, **kwargs)
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(
|
||||
{
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import os
|
||||
import inspect
|
||||
from typing import Dict
|
||||
|
||||
import folder_paths
|
||||
@@ -7,7 +6,7 @@ import folder_paths
|
||||
from ..utils import load_module
|
||||
|
||||
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
||||
efficieny_dir_names = ["ImpactPack", "ComfyUI-Impact-Pack"]
|
||||
efficieny_dir_names = ["ImpactPack", "ComfyUI-Impact-Pack", "comfyui-impact-pack"]
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
@@ -36,6 +35,9 @@ try:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
inputs = FaceDetailer.INPUT_TYPES()
|
||||
if not "optional" in inputs:
|
||||
inputs["optional"] = {}
|
||||
|
||||
inputs["optional"]["enabled"] = (
|
||||
"BOOLEAN",
|
||||
{"default": True, "label_on": "enabled", "label_off": "disabled"},
|
||||
@@ -75,6 +77,9 @@ try:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
inputs = FaceDetailerPipe.INPUT_TYPES()
|
||||
if not "optional" in inputs:
|
||||
inputs["optional"] = {}
|
||||
|
||||
inputs["optional"]["enabled"] = (
|
||||
"BOOLEAN",
|
||||
{"default": True, "label_on": "enabled", "label_off": "disabled"},
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
import os
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import logging
|
||||
|
||||
import folder_paths
|
||||
import comfy.model_management as model_management
|
||||
@@ -9,12 +10,24 @@ import comfy.model_management as model_management
|
||||
from ...model_utils import download_file
|
||||
|
||||
|
||||
lama = None
|
||||
model_dir = os.path.join(folder_paths.models_dir, "lama")
|
||||
|
||||
if "lama" not in folder_paths.folder_names_and_paths:
|
||||
folder_paths.folder_names_and_paths["lama"] = ([model_dir], folder_paths.supported_pt_extensions)
|
||||
|
||||
gpu = model_management.get_torch_device()
|
||||
cpu = torch.device("cpu")
|
||||
model_dir = os.path.join(folder_paths.models_dir, "lama")
|
||||
model_url = "https://github.com/Sanster/models/releases/download/add_big_lama/big-lama.pt"
|
||||
model_sha = "344c77bbcb158f17dd143070d1e789f38a66c04202311ae3a258ef66667a9ea9"
|
||||
|
||||
_models = {
|
||||
"big-lama.pt": {
|
||||
"url": "https://github.com/Sanster/models/releases/download/add_big_lama/big-lama.pt",
|
||||
"sha": "344c77bbcb158f17dd143070d1e789f38a66c04202311ae3a258ef66667a9ea9",
|
||||
},
|
||||
"anime-manga-big-lama.pt": {
|
||||
"url": "https://github.com/Sanster/models/releases/download/AnimeMangaInpainting/anime-manga-big-lama.pt",
|
||||
"sha": "479d3afdcb7ed2fd944ed4ebcc39ca45b33491f0f2e43eb1000bd623cfb41823",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def ceil_modulo(x, mod):
|
||||
@@ -30,19 +43,43 @@ def pad_tensor_to_modulo(img, mod):
|
||||
return F.pad(img, pad=(0, out_width - width, 0, out_height - height), mode="reflect")
|
||||
|
||||
|
||||
def load_model():
|
||||
global lama
|
||||
if lama is None:
|
||||
model_path = os.path.join(model_dir, "big-lama.pt")
|
||||
download_file(model_url, model_path, model_sha)
|
||||
class LoadLaMaModel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"model_name": (
|
||||
list(set(["big-lama.pt", "anime-manga-big-lama.pt"] + folder_paths.get_filename_list("lama"))),
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LAMA",)
|
||||
CATEGORY = "ArtVenture/Inpainting"
|
||||
FUNCTION = "load_model"
|
||||
|
||||
def load_model(self, model_name: str):
|
||||
model_path = folder_paths.get_full_path("lama", model_name)
|
||||
if model_path is None:
|
||||
if model_name in _models:
|
||||
model_url = _models[model_name]["url"]
|
||||
model_sha = _models[model_name]["sha"]
|
||||
logging.info(f"Downloading {model_name} into {model_dir}")
|
||||
model_path = os.path.join(model_dir, model_name)
|
||||
download_file(model_url, model_path, model_sha)
|
||||
else:
|
||||
raise Exception(f"Not found model {model_name}")
|
||||
|
||||
lama = torch.jit.load(model_path, map_location="cpu")
|
||||
lama.eval()
|
||||
|
||||
return lama
|
||||
return (lama,)
|
||||
|
||||
|
||||
class LaMaInpaint:
|
||||
class LaMaInpaint(LoadLaMaModel):
|
||||
def __init__(self):
|
||||
self.model_name = "big-lama.pt"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
@@ -50,25 +87,22 @@ class LaMaInpaint:
|
||||
"image": ("IMAGE",),
|
||||
"mask": ("MASK",),
|
||||
},
|
||||
"optional": {"device_mode": (["AUTO", "Prefer GPU", "CPU"],)},
|
||||
"optional": {
|
||||
"device_mode": (["AUTO", "Prefer GPU", "CPU"],),
|
||||
"lama_model": ("LAMA",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Inpainting"
|
||||
FUNCTION = "lama_inpaint"
|
||||
|
||||
def lama_inpaint(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
device_mode="AUTO",
|
||||
):
|
||||
def lama_inpaint(self, image: torch.Tensor, mask: torch.Tensor, device_mode="AUTO", lama_model=None):
|
||||
if image.shape[0] != mask.shape[0]:
|
||||
raise Exception("Image and mask must have the same batch size")
|
||||
|
||||
device = gpu if device_mode != "CPU" else cpu
|
||||
|
||||
model = load_model()
|
||||
model = lama_model or self.load_model(self.model_name)
|
||||
model.to(device)
|
||||
|
||||
try:
|
||||
|
||||
+82
-37
@@ -5,10 +5,10 @@ from PIL import Image, ImageOps
|
||||
from typing import Dict
|
||||
|
||||
from .sam.nodes import SAMLoader, GetSAMEmbedding, SAMEmbeddingToImage
|
||||
from .lama import LaMaInpaint
|
||||
from .lama import LoadLaMaModel, LaMaInpaint
|
||||
|
||||
from ..masking import get_crop_region, expand_crop_region
|
||||
from ..image_utils import ResizeMode, resize_image, flatten_image
|
||||
from ..image_utils import ResizeMode, resize_image
|
||||
from ..utils import numpy2pil, tensor2pil, pil2tensor
|
||||
|
||||
|
||||
@@ -21,46 +21,54 @@ class PrepareImageAndMaskForInpaint:
|
||||
"mask": ("MASK",),
|
||||
"mask_blur": ("INT", {"default": 4, "min": 0, "max": 64}),
|
||||
"inpaint_masked": ("BOOLEAN", {"default": False}),
|
||||
"mask_padding": ("INT", {"default": 32, "min": 0, "max": 256}),
|
||||
"mask_padding": ("INT", {"default": 32, "min": 0, "max": 1024}),
|
||||
"width": ("INT", {"default": 0, "min": 0, "max": 2048}),
|
||||
"height": ("INT", {"default": 0, "min": 0, "max": 2048}),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"controlnet_image": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "CROP_REGION")
|
||||
RETURN_NAMES = ("inpaint_image", "inpaint_mask", "overlay_image", "crop_region")
|
||||
CATEGORY = "Art Venture/Inpainting"
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "IMAGE", "CROP_REGION", "IMAGE")
|
||||
RETURN_NAMES = ("inpaint_image", "inpaint_mask", "overlay_image", "crop_region", "controlnet_image")
|
||||
CATEGORY = "ArtVenture/Inpainting"
|
||||
FUNCTION = "prepare"
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
mask: torch.Tensor,
|
||||
# resize_mode: str,
|
||||
mask_blur: int,
|
||||
inpaint_masked: bool,
|
||||
mask_padding: int,
|
||||
width: int,
|
||||
height: int,
|
||||
controlnet_image: torch.Tensor = None,
|
||||
):
|
||||
if image.shape[0] != mask.shape[0]:
|
||||
raise ValueError("image and mask must have same batch size")
|
||||
|
||||
if controlnet_image is not None and image.shape[0] != controlnet_image.shape[0]:
|
||||
raise ValueError("image and controlnet_image must have same batch size")
|
||||
|
||||
if image.shape[1] != mask.shape[1] or image.shape[2] != mask.shape[2]:
|
||||
raise ValueError("image and mask must have same dimensions")
|
||||
|
||||
if width == 0 and height == 0:
|
||||
height, width = image.shape[1:3]
|
||||
|
||||
sourceheight, sourcewidth = image.shape[1:3]
|
||||
# These are only used if inpaint_masked is True
|
||||
out_width, out_height = width, height
|
||||
if inpaint_masked and out_width == 0 and out_height == 0:
|
||||
out_height, out_width = image.shape[1:3]
|
||||
|
||||
source_height, source_width = image.shape[1:3]
|
||||
|
||||
masks = []
|
||||
images = []
|
||||
overlay_masks = []
|
||||
masks = []
|
||||
overlay_images = []
|
||||
crop_regions = []
|
||||
processed_controlnet_images = []
|
||||
|
||||
for img, msk in zip(image, mask):
|
||||
for idx, (img, msk) in enumerate(zip(image, mask)):
|
||||
np_mask: np.ndarray = msk.cpu().numpy()
|
||||
|
||||
if mask_blur > 0:
|
||||
@@ -68,41 +76,76 @@ class PrepareImageAndMaskForInpaint:
|
||||
np_mask = cv2.GaussianBlur(np_mask, (kernel_size, kernel_size), mask_blur)
|
||||
|
||||
pil_mask = numpy2pil(np_mask, "L")
|
||||
crop_region = None
|
||||
pil_img = tensor2pil(img)
|
||||
|
||||
# --- LOGIC SEPARATION ---
|
||||
|
||||
if inpaint_masked:
|
||||
# --- MODE 1: CROP AND RESIZE ---
|
||||
crop_region = get_crop_region(np_mask, mask_padding)
|
||||
crop_region = expand_crop_region(crop_region, width, height, sourcewidth, sourceheight)
|
||||
# crop mask
|
||||
overlay_mask = pil_mask
|
||||
pil_mask = resize_image(pil_mask.crop(crop_region), width, height, ResizeMode.RESIZE_TO_FIT)
|
||||
pil_mask = pil_mask.convert("L")
|
||||
crop_region = expand_crop_region(crop_region, out_width, out_height, source_width, source_height)
|
||||
|
||||
cropped_img = pil_img.crop(crop_region)
|
||||
cropped_mask = pil_mask.crop(crop_region)
|
||||
|
||||
final_pil_img = resize_image(cropped_img, out_width, out_height, ResizeMode.RESIZE_TO_FIT)
|
||||
final_pil_mask = resize_image(cropped_mask, out_width, out_height, ResizeMode.RESIZE_TO_FIT).convert(
|
||||
"L"
|
||||
)
|
||||
|
||||
if controlnet_image is not None:
|
||||
pil_cimg = tensor2pil(controlnet_image[idx])
|
||||
cn_source_width, cn_source_height = pil_cimg.size
|
||||
scale_x = cn_source_width / source_width
|
||||
scale_y = cn_source_height / source_height
|
||||
|
||||
cn_target_width = int(out_width * scale_x)
|
||||
cn_target_height = int(out_height * scale_y)
|
||||
|
||||
x1, y1, x2, y2 = crop_region
|
||||
cn_crop_region = (int(x1 * scale_x), int(y1 * scale_y), int(x2 * scale_x), int(y2 * scale_y))
|
||||
cropped_cn_img = pil_cimg.crop(cn_crop_region)
|
||||
final_cn_img = resize_image(
|
||||
cropped_cn_img, cn_target_width, cn_target_height, ResizeMode.RESIZE_TO_FIT
|
||||
)
|
||||
processed_controlnet_images.append(pil2tensor(final_cn_img))
|
||||
|
||||
else:
|
||||
np_mask = np.clip((np_mask.astype(np.float32)) * 2, 0, 255).astype(np.uint8)
|
||||
overlay_mask = numpy2pil(np_mask, "L")
|
||||
# --- MODE 2: PASS-THROUGH (NO RESIZING) ---
|
||||
final_pil_img = pil_img
|
||||
final_pil_mask = pil_mask # Already blurred if requested
|
||||
crop_region = (0, 0, source_width, source_height)
|
||||
|
||||
pil_img = tensor2pil(img)
|
||||
pil_img = flatten_image(pil_img)
|
||||
if controlnet_image is not None:
|
||||
# Simply pass the original controlnet image through
|
||||
final_cn_img = tensor2pil(controlnet_image[idx])
|
||||
processed_controlnet_images.append(pil2tensor(final_cn_img))
|
||||
|
||||
# --- COMMON LOGIC FOR BOTH MODES ---
|
||||
|
||||
# The overlay/preview should always be based on the original full-size image
|
||||
image_masked = Image.new("RGBa", (pil_img.width, pil_img.height))
|
||||
image_masked.paste(pil_img.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(overlay_mask))
|
||||
# The mask used here is the potentially blurred one, but before any cropping/resizing
|
||||
image_masked.paste(pil_img.convert("RGBA").convert("RGBa"), mask=ImageOps.invert(pil_mask))
|
||||
overlay_images.append(pil2tensor(image_masked.convert("RGBA")))
|
||||
overlay_masks.append(pil2tensor(overlay_mask))
|
||||
|
||||
if crop_region is not None:
|
||||
pil_img = resize_image(pil_img.crop(crop_region), width, height, ResizeMode.RESIZE_TO_FIT)
|
||||
else:
|
||||
crop_region = (0, 0, 0, 0)
|
||||
|
||||
images.append(pil2tensor(pil_img))
|
||||
masks.append(pil2tensor(pil_mask))
|
||||
images.append(pil2tensor(final_pil_img))
|
||||
masks.append(pil2tensor(final_pil_mask))
|
||||
crop_regions.append(torch.tensor(crop_region, dtype=torch.int64))
|
||||
|
||||
if processed_controlnet_images:
|
||||
final_controlnet_tensor = torch.cat(processed_controlnet_images, dim=0)
|
||||
else:
|
||||
# If no controlnet image is provided, create a black 64x64 placeholder
|
||||
batch_size = image.shape[0]
|
||||
final_controlnet_tensor = torch.zeros((batch_size, 64, 64, 3), dtype=torch.float32, device=image.device)
|
||||
|
||||
return (
|
||||
torch.cat(images, dim=0),
|
||||
torch.cat(masks, dim=0),
|
||||
torch.cat(overlay_images, dim=0),
|
||||
torch.stack(crop_regions),
|
||||
torch.stack(crop_regions, dim=0),
|
||||
final_controlnet_tensor,
|
||||
)
|
||||
|
||||
|
||||
@@ -118,7 +161,7 @@ class OverlayInpaintedLatent:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LATENT",)
|
||||
CATEGORY = "Art Venture/Inpainting"
|
||||
CATEGORY = "ArtVenture/Inpainting"
|
||||
FUNCTION = "overlay"
|
||||
|
||||
def overlay(self, original: Dict, inpainted: Dict, mask: torch.Tensor):
|
||||
@@ -162,7 +205,7 @@ class OverlayInpaintedImage:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Inpainting"
|
||||
CATEGORY = "ArtVenture/Inpainting"
|
||||
FUNCTION = "overlay"
|
||||
|
||||
def overlay(self, inpainted: torch.Tensor, overlay_image: torch.Tensor, crop_region: torch.Tensor):
|
||||
@@ -198,6 +241,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"AV_SAMLoader": SAMLoader,
|
||||
"GetSAMEmbedding": GetSAMEmbedding,
|
||||
"SAMEmbeddingToImage": SAMEmbeddingToImage,
|
||||
"LoadLaMaModel": LoadLaMaModel,
|
||||
"LaMaInpaint": LaMaInpaint,
|
||||
"PrepareImageAndMaskForInpaint": PrepareImageAndMaskForInpaint,
|
||||
"OverlayInpaintedLatent": OverlayInpaintedLatent,
|
||||
@@ -208,6 +252,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AV_SAMLoader": "SAM Loader",
|
||||
"GetSAMEmbedding": "Get SAM Embedding",
|
||||
"SAMEmbeddingToImage": "SAM Embedding to Image",
|
||||
"LoadLaMaModel": "LaMa Loader",
|
||||
"LaMaInpaint": "LaMa Remove Object",
|
||||
"PrepareImageAndMaskForInpaint": "Prepare Image & Mask for Inpaint",
|
||||
"OverlayInpaintedLatent": "Overlay Inpainted Latent",
|
||||
|
||||
@@ -9,12 +9,13 @@ import comfy.utils
|
||||
|
||||
from ...utils import ensure_package, tensor2pil, pil2tensor
|
||||
|
||||
folder_paths.folder_names_and_paths["sams"] = (
|
||||
[
|
||||
os.path.join(folder_paths.models_dir, "sams"),
|
||||
],
|
||||
folder_paths.supported_pt_extensions,
|
||||
)
|
||||
if "sams" not in folder_paths.folder_names_and_paths:
|
||||
folder_paths.folder_names_and_paths["sams"] = (
|
||||
[
|
||||
os.path.join(folder_paths.models_dir, "sams"),
|
||||
],
|
||||
folder_paths.supported_pt_extensions,
|
||||
)
|
||||
|
||||
gpu = model_management.get_torch_device()
|
||||
cpu = torch.device("cpu")
|
||||
@@ -32,7 +33,7 @@ class SAMLoader:
|
||||
RETURN_TYPES = ("AV_SAM_MODEL",)
|
||||
RETURN_NAMES = ("sam_model",)
|
||||
FUNCTION = "load_model"
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
CATEGORY = "ArtVenture/Segmentation"
|
||||
|
||||
def load_model(self, model_name):
|
||||
modelname = folder_paths.get_full_path("sams", model_name)
|
||||
@@ -68,7 +69,7 @@ class GetSAMEmbedding:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("SAM_EMBEDDING",)
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
CATEGORY = "ArtVenture/Segmentation"
|
||||
FUNCTION = "get_sam_embedding"
|
||||
|
||||
def get_sam_embedding(self, image, sam_model, device_mode="AUTO"):
|
||||
@@ -102,7 +103,7 @@ class SAMEmbeddingToImage:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
CATEGORY = "ArtVenture/Segmentation"
|
||||
FUNCTION = "sam_embedding_to_noise_image"
|
||||
|
||||
def sam_embedding_to_noise_image(self, embedding: np.ndarray):
|
||||
|
||||
@@ -122,7 +122,7 @@ class BlipLoader:
|
||||
|
||||
RETURN_TYPES = ("BLIP_MODEL",)
|
||||
FUNCTION = "load_blip"
|
||||
CATEGORY = "Art Venture/Captioning"
|
||||
CATEGORY = "ArtVenture/Captioning"
|
||||
|
||||
def load_blip(self, model_name):
|
||||
return (load_blip(model_name),)
|
||||
@@ -139,7 +139,7 @@ class DownloadAndLoadBlip:
|
||||
|
||||
RETURN_TYPES = ("BLIP_MODEL",)
|
||||
FUNCTION = "download_and_load_blip"
|
||||
CATEGORY = "Art Venture/Captioning"
|
||||
CATEGORY = "ArtVenture/Captioning"
|
||||
|
||||
def download_and_load_blip(self, model_name):
|
||||
if model_name not in folder_paths.get_filename_list("blip"):
|
||||
@@ -191,7 +191,7 @@ class BlipCaption:
|
||||
RETURN_NAMES = ("caption",)
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
FUNCTION = "blip_caption"
|
||||
CATEGORY = "Art Venture/Captioning"
|
||||
CATEGORY = "ArtVenture/Captioning"
|
||||
|
||||
def blip_caption(
|
||||
self, image, min_length, max_length, device_mode="AUTO", prefix="", suffix="", enabled=True, blip_model=None
|
||||
|
||||
@@ -79,7 +79,7 @@ class DeepDanbooruCaption:
|
||||
RETURN_NAMES = ("caption",)
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
FUNCTION = "caption"
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
|
||||
def caption(
|
||||
self,
|
||||
|
||||
@@ -25,10 +25,10 @@ from transformers.modeling_outputs import (
|
||||
)
|
||||
from transformers.modeling_utils import (
|
||||
PreTrainedModel,
|
||||
apply_chunking_to_forward,
|
||||
find_pruneable_heads_and_indices,
|
||||
prune_linear_layer,
|
||||
)
|
||||
from transformers.pytorch_utils import apply_chunking_to_forward
|
||||
from transformers.utils import logging
|
||||
from transformers.models.bert.configuration_bert import BertConfig
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import comfy.model_management
|
||||
from .utils import load_module
|
||||
|
||||
custom_nodes = folder_paths.get_folder_paths("custom_nodes")
|
||||
ip_adapter_dir_names = ["IPAdapter", "ComfyUI_IPAdapter_plus"]
|
||||
ip_adapter_dir_names = ["IPAdapter", "ComfyUI_IPAdapter_plus", "comfyui_ipadapter_plus"]
|
||||
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
@@ -50,7 +50,7 @@ try:
|
||||
|
||||
RETURN_TYPES = ("IPADAPTER",)
|
||||
RETURN_NAMES = ("pipeline",)
|
||||
CATEGORY = "Art Venture/IP Adapter"
|
||||
CATEGORY = "ArtVenture/IP Adapter"
|
||||
FUNCTION = "load_ip_adapter"
|
||||
|
||||
def load_ip_adapter(self, ip_adapter_name, clip_name):
|
||||
@@ -88,7 +88,7 @@ try:
|
||||
|
||||
RETURN_TYPES = ("MODEL", "IPADAPTER", "CLIP_VISION")
|
||||
RETURN_NAMES = ("model", "pipeline", "clip_vision")
|
||||
CATEGORY = "Art Venture/IP Adapter"
|
||||
CATEGORY = "ArtVenture/IP Adapter"
|
||||
FUNCTION = "apply_ip_adapter"
|
||||
|
||||
def apply_ip_adapter(
|
||||
|
||||
@@ -152,7 +152,7 @@ class ISNetLoader:
|
||||
|
||||
RETURN_TYPES = ("ISNET_MODEL",)
|
||||
FUNCTION = "load_isnet"
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
CATEGORY = "ArtVenture/Segmentation"
|
||||
|
||||
def load_isnet(self, model_name):
|
||||
return (load_isnet_model(model_name),)
|
||||
@@ -169,7 +169,7 @@ class DownloadISNetModel:
|
||||
|
||||
RETURN_TYPES = ("ISNET_MODEL",)
|
||||
FUNCTION = "download_isnet"
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
CATEGORY = "ArtVenture/Segmentation"
|
||||
|
||||
def download_isnet(self, model_name):
|
||||
if model_name not in folder_paths.get_filename_list("isnet"):
|
||||
@@ -200,7 +200,7 @@ class ISNetSegment:
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("segmented", "mask")
|
||||
CATEGORY = "Art Venture/Segmentation"
|
||||
CATEGORY = "ArtVenture/Segmentation"
|
||||
FUNCTION = "segment_isnet"
|
||||
|
||||
def segment_isnet(self, images: torch.Tensor, threshold, device_mode="AUTO", enabled=True, isnet_model=None):
|
||||
|
||||
+565
-116
@@ -1,38 +1,114 @@
|
||||
import os
|
||||
import base64
|
||||
import json
|
||||
import requests
|
||||
import os
|
||||
from enum import Enum
|
||||
from io import BytesIO
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import requests
|
||||
import torch
|
||||
from PIL import Image
|
||||
from pydantic import BaseModel
|
||||
from torch import Tensor
|
||||
from pydantic import BaseModel
|
||||
from typing import List, Dict, Union, Optional, Any
|
||||
|
||||
from ..utils import ensure_package, pil2base64, pil2tensor, tensor2pil
|
||||
|
||||
|
||||
def image_urls_to_tensor(
|
||||
image_urls: List[str], timeout: int | None = 60
|
||||
) -> Optional[Tensor]:
|
||||
tensors: List[Tensor] = []
|
||||
|
||||
for url in image_urls:
|
||||
try:
|
||||
if isinstance(url, str) and url.startswith("data:"):
|
||||
comma_index = url.find(",")
|
||||
if comma_index == -1:
|
||||
continue
|
||||
header = url[:comma_index]
|
||||
data_part = url[comma_index + 1 :]
|
||||
|
||||
if ";base64" in header:
|
||||
raw_bytes = base64.b64decode(data_part)
|
||||
else:
|
||||
raw_bytes = data_part.encode("utf-8")
|
||||
|
||||
image = Image.open(BytesIO(raw_bytes)).convert("RGB")
|
||||
else:
|
||||
resp = requests.get(url, timeout=timeout)
|
||||
resp.raise_for_status()
|
||||
image = Image.open(BytesIO(resp.content)).convert("RGB")
|
||||
|
||||
tensor = pil2tensor(image)
|
||||
tensors.append(tensor)
|
||||
except Exception:
|
||||
# Silently skip any image that fails to load/parse
|
||||
continue
|
||||
|
||||
if tensors:
|
||||
return torch.cat(tensors, dim=0)
|
||||
|
||||
return None
|
||||
|
||||
from ..utils import ensure_package, tensor2pil, pil2base64
|
||||
|
||||
gpt_models = [
|
||||
"gpt-3.5-turbo",
|
||||
"gpt-3.5-turbo-16k",
|
||||
"gpt-5.4-mini",
|
||||
"gpt-5.4",
|
||||
"gpt-5",
|
||||
"gpt-5-mini",
|
||||
"gpt-5-nano",
|
||||
"gpt-5-chat-latest",
|
||||
"gpt-4o",
|
||||
"gpt-4o-mini",
|
||||
"gpt-4.1",
|
||||
"gpt-4.1-mini",
|
||||
"gpt-4.1-nano",
|
||||
"gpt-4-turbo",
|
||||
"gpt-4-vision-preview",
|
||||
"gpt-4-turbo-preview",
|
||||
"gpt-4-0125-preview",
|
||||
"gpt-4-1106-preview",
|
||||
"gpt-4-0613",
|
||||
"gpt-4",
|
||||
"gpt-4-vision-preview",
|
||||
"o1",
|
||||
"o1-mini",
|
||||
"o1-preview",
|
||||
"o1-pro",
|
||||
"o3",
|
||||
"o3-mini",
|
||||
"o3-pro",
|
||||
"o4-mini",
|
||||
]
|
||||
|
||||
gpt_vision_models = ["gpt-4o", "gpt-4o-mini", "gpt-4-turbo", "gpt-4-turbo-preview", "gpt-4-vision-preview"]
|
||||
|
||||
claude3_models = [
|
||||
claude_models = [
|
||||
"claude-sonnet-4-5-20250929",
|
||||
"claude-haiku-4-5-20251001",
|
||||
"claude-sonnet-4-20250514",
|
||||
"claude-opus-4-1-20250805",
|
||||
"claude-opus-4-20250514",
|
||||
"claude-3-7-sonnet-latest",
|
||||
"claude-3-7-sonnet-20250219",
|
||||
"claude-3-5-sonnet-latest",
|
||||
"claude-3-5-sonnet-20241022",
|
||||
"claude-3-5-haiku-20241022",
|
||||
"claude-3-opus-latest",
|
||||
"claude-3-opus-20240229",
|
||||
"claude-3-sonnet-20240229",
|
||||
"claude-3-haiku-20240307",
|
||||
]
|
||||
claude2_models = ["claude-2.1"]
|
||||
|
||||
nano_banana_models = {
|
||||
"nano-banana": "gemini-2.5-flash-image",
|
||||
"nano-banana-2": "gemini-3.1-flash-image-preview",
|
||||
"nano-banana-pro": "gemini-3-pro-image-preview",
|
||||
}
|
||||
|
||||
gemini_models = [
|
||||
"gemini-3-flash-preview",
|
||||
"gemini-3.1-pro-preview",
|
||||
"gemini-2.5-flash",
|
||||
"gemini-2.5-flash-lite",
|
||||
"gemini-2.5-pro",
|
||||
"gemini-2.0-flash",
|
||||
"gemini-2.0-flash-lite",
|
||||
]
|
||||
|
||||
aws_regions = [
|
||||
"us-east-1",
|
||||
@@ -48,24 +124,31 @@ aws_regions = [
|
||||
|
||||
bedrock_anthropic_versions = ["bedrock-2023-05-31"]
|
||||
|
||||
bedrock_claude3_models = [
|
||||
bedrock_claude_models = [
|
||||
"anthropic.claude-opus-4-20250514-v1:0",
|
||||
"anthropic.claude-sonnet-4-20250514-v1:0",
|
||||
"anthropic.claude-3-7-sonnet-20250219-v1:0",
|
||||
"anthropic.claude-3-5-sonnet-20241022-v2:0",
|
||||
"anthropic.claude-3-5-haiku-20241022-v1:0",
|
||||
"anthropic.claude-3-haiku-20240307-v1:0",
|
||||
"anthropic.claude-3-sonnet-20240229-v1:0",
|
||||
"anthropic.claude-3-opus-20240229-v1:0",
|
||||
]
|
||||
|
||||
bedrock_claude2_models = [
|
||||
"anthropic.claude-v2",
|
||||
"anthropic.claude-v2.1",
|
||||
]
|
||||
|
||||
bedrock_mistral_models = [
|
||||
"mistral.mistral-7b-instruct-v0:2",
|
||||
"mistral.mixtral-8x7b-instruct-v0:1",
|
||||
"mistral.mistral-large-2402-v1:0",
|
||||
]
|
||||
|
||||
all_models = (
|
||||
gpt_models
|
||||
+ claude_models
|
||||
+ gemini_models
|
||||
+ bedrock_claude_models
|
||||
+ bedrock_mistral_models
|
||||
)
|
||||
|
||||
default_system_prompt = "You are a useful AI agent."
|
||||
|
||||
|
||||
@@ -75,6 +158,12 @@ class LLMConfig(BaseModel):
|
||||
temperature: float
|
||||
|
||||
|
||||
class NanoBananaConfig(LLMConfig):
|
||||
modalities: str = "image+text"
|
||||
aspect_ratio: str = "auto"
|
||||
resolution: str = "1K"
|
||||
|
||||
|
||||
class LLMMessageRole(str, Enum):
|
||||
system = "system"
|
||||
user = "user"
|
||||
@@ -84,13 +173,19 @@ class LLMMessageRole(str, Enum):
|
||||
class LLMMessage(BaseModel):
|
||||
role: LLMMessageRole = LLMMessageRole.user
|
||||
text: str
|
||||
image: Optional[str] = None # base64 enoded image
|
||||
images: Optional[List[str]] = None # list of base64 encoded images
|
||||
|
||||
def to_openai_message(self):
|
||||
content = [{"type": "text", "text": self.text}]
|
||||
content: List[Dict[str, Any]] = [{"type": "text", "text": self.text}]
|
||||
|
||||
if self.image:
|
||||
content.insert(0, {"type": "image_url", "image_url": {"url": f"data:image/png;base64,{self.text}"}})
|
||||
if self.images:
|
||||
for img in self.images:
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": f"data:image/png;base64,{img}"},
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"role": self.role,
|
||||
@@ -98,22 +193,48 @@ class LLMMessage(BaseModel):
|
||||
}
|
||||
|
||||
def to_claude_message(self):
|
||||
content = [{"type": "text", "text": self.text}]
|
||||
content: List[Dict[str, Any]] = [{"type": "text", "text": self.text}]
|
||||
|
||||
if self.image:
|
||||
content.insert(
|
||||
0,
|
||||
{
|
||||
"type": "image",
|
||||
"source": {"type": "base64", "media_type": "image/png", "data": self.image},
|
||||
},
|
||||
)
|
||||
if self.images:
|
||||
for img in reversed(self.images):
|
||||
content.append(
|
||||
{
|
||||
"type": "image",
|
||||
"source": {
|
||||
"type": "base64",
|
||||
"media_type": "image/png",
|
||||
"data": img,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
return {
|
||||
"role": self.role,
|
||||
"content": content,
|
||||
}
|
||||
|
||||
def to_gemini_message(self):
|
||||
parts: List[Dict[str, Any]] = [{"text": self.text}]
|
||||
|
||||
if self.images:
|
||||
for img in self.images:
|
||||
parts.append(
|
||||
{
|
||||
"inline_data": {
|
||||
"mime_type": "image/png",
|
||||
"data": img,
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
# Gemini uses "model" and "user" roles instead of "assistant" and "user"
|
||||
role = "model" if self.role == "assistant" else "user"
|
||||
|
||||
return {
|
||||
"role": role,
|
||||
"parts": parts,
|
||||
}
|
||||
|
||||
|
||||
class OpenAIApi(BaseModel):
|
||||
api_key: str
|
||||
@@ -121,7 +242,7 @@ class OpenAIApi(BaseModel):
|
||||
timeout: Optional[int] = 60
|
||||
|
||||
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
|
||||
if config.model not in gpt_models:
|
||||
if config.model in all_models and config.model not in gpt_models:
|
||||
raise Exception(f"Must provide an OpenAI model, got {config.model}")
|
||||
|
||||
formated_messages = [m.to_openai_message() for m in messages]
|
||||
@@ -132,7 +253,6 @@ class OpenAIApi(BaseModel):
|
||||
"model": config.model,
|
||||
"max_tokens": config.max_token,
|
||||
"temperature": config.temperature,
|
||||
# "seed": seed,
|
||||
}
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
|
||||
@@ -140,9 +260,81 @@ class OpenAIApi(BaseModel):
|
||||
data: Dict = response.json()
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
raise Exception(data.get("error").get("message"))
|
||||
return (f"OpenAI API error: {data.get('error').get('message')}", None)
|
||||
|
||||
return data["choices"][0]["message"]["content"]
|
||||
text = data["choices"][0]["message"]["content"]
|
||||
return (text, None)
|
||||
|
||||
def complete(self, prompt: str, config: LLMConfig, seed=None):
|
||||
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
|
||||
|
||||
return self.chat(messages, config, seed)
|
||||
|
||||
|
||||
class OpenRouterApi(BaseModel):
|
||||
api_key: str
|
||||
endpoint: Optional[str] = "https://openrouter.ai/api/v1"
|
||||
timeout: Optional[int] = 60
|
||||
|
||||
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
|
||||
formated_messages = [m.to_openai_message() for m in messages]
|
||||
|
||||
url = f"{self.endpoint}/chat/completions"
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
|
||||
data = {
|
||||
"messages": formated_messages,
|
||||
"model": config.model,
|
||||
"max_tokens": config.max_token,
|
||||
"temperature": config.temperature,
|
||||
}
|
||||
|
||||
if isinstance(config, NanoBananaConfig):
|
||||
modalities = (config.modalities or "text+image").split("+")
|
||||
modalities = [modality.strip().lower() for modality in modalities]
|
||||
aspectRatio = (
|
||||
config.aspect_ratio
|
||||
if (config.aspect_ratio and config.aspect_ratio != "auto")
|
||||
else None
|
||||
)
|
||||
|
||||
data["modalities"] = modalities
|
||||
data["image_config"] = {}
|
||||
|
||||
if aspectRatio is not None:
|
||||
data["image_config"]["aspect_ratio"] = aspectRatio
|
||||
|
||||
# nano-banana 1 does not support imageSize
|
||||
if "nano-banana-" in config.model:
|
||||
data["image_config"]["image_size"] = config.resolution
|
||||
|
||||
model = nano_banana_models[config.model]
|
||||
data["model"] = f"google/{model}"
|
||||
|
||||
response = requests.post(url, json=data, headers=headers, timeout=self.timeout)
|
||||
data: Dict = response.json()
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
return (f"OpenRouter API error: {data.get('error').get('message')}", None)
|
||||
|
||||
message = data["choices"][0]["message"]
|
||||
text = message["content"]
|
||||
images = None
|
||||
|
||||
if message.get("images", None) is not None:
|
||||
urls: List[str] = []
|
||||
for m in message["images"]:
|
||||
if isinstance(m, dict):
|
||||
if (
|
||||
"image_url" in m
|
||||
and isinstance(m["image_url"], dict)
|
||||
and "url" in m["image_url"]
|
||||
):
|
||||
urls.append(m["image_url"]["url"])
|
||||
|
||||
images = image_urls_to_tensor(urls, timeout=self.timeout)
|
||||
|
||||
return (text, images)
|
||||
|
||||
def complete(self, prompt: str, config: LLMConfig, seed=None):
|
||||
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
|
||||
@@ -157,8 +349,8 @@ class ClaudeApi(BaseModel):
|
||||
timeout: Optional[int] = 60
|
||||
|
||||
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
|
||||
if config.model not in claude3_models:
|
||||
raise Exception(f"Must provide a Claude v3 model, got {config.model}")
|
||||
if config.model in all_models and config.model not in claude_models:
|
||||
raise Exception(f"Must provide a Claude model, got {config.model}")
|
||||
|
||||
system_message = [m for m in messages if m.role == "system"]
|
||||
user_messages = [m for m in messages if m.role != "system"]
|
||||
@@ -178,30 +370,114 @@ class ClaudeApi(BaseModel):
|
||||
data: Dict = response.json()
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
raise Exception(data.get("error").get("message"))
|
||||
return (data.get("error").get("message"), None)
|
||||
|
||||
return data["content"][0]["text"]
|
||||
text = data["content"][0]["text"]
|
||||
return (text, None)
|
||||
|
||||
def complete(self, prompt: str, config: LLMConfig, seed=None):
|
||||
if config.model not in claude2_models:
|
||||
raise Exception(f"Must provide a Claude v2 model, got {config.model}")
|
||||
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
|
||||
return self.chat(messages, config, seed)
|
||||
|
||||
|
||||
class GeminiApi(BaseModel):
|
||||
api_key: str
|
||||
endpoint: Optional[str] = "https://generativelanguage.googleapis.com/v1beta"
|
||||
timeout: Optional[int] = 60
|
||||
|
||||
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
|
||||
if config.model in all_models and config.model not in gemini_models:
|
||||
raise Exception(f"Must provide a Gemini model, got {config.model}")
|
||||
|
||||
system_message = [m for m in messages if m.role == "system"]
|
||||
user_messages = [m for m in messages if m.role != "system"]
|
||||
|
||||
if not user_messages:
|
||||
return (
|
||||
"Gemini API error: At least one user message is required. System messages alone are not sufficient.",
|
||||
None,
|
||||
)
|
||||
|
||||
formated_messages = [m.to_gemini_message() for m in user_messages]
|
||||
|
||||
url = f"{self.endpoint}/models/{config.model}:generateContent"
|
||||
headers = {"x-goog-api-key": self.api_key}
|
||||
|
||||
prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
|
||||
url = f"{self.endpoint}/complete"
|
||||
data = {
|
||||
"prompt": prompt,
|
||||
"max_tokens_to_sample": config.max_token,
|
||||
"temperature": config.temperature,
|
||||
"contents": formated_messages,
|
||||
"generationConfig": {
|
||||
"maxOutputTokens": config.max_token,
|
||||
"temperature": config.temperature,
|
||||
},
|
||||
}
|
||||
headers = {"x-api-key": self.api_key, "anthropic-version": self.version}
|
||||
|
||||
response = requests.post(url, json=data, headers=headers, timeout=self.timeout)
|
||||
if isinstance(config, NanoBananaConfig):
|
||||
modalities = (config.modalities or "text+image").split("+")
|
||||
modalities = [modality.strip().capitalize() for modality in modalities]
|
||||
aspectRatio = (
|
||||
config.aspect_ratio
|
||||
if (config.aspect_ratio and config.aspect_ratio != "auto")
|
||||
else None
|
||||
)
|
||||
|
||||
data["generationConfig"]["responseModalities"] = modalities
|
||||
data["generationConfig"]["imageConfig"] = {
|
||||
"aspectRatio": aspectRatio,
|
||||
}
|
||||
|
||||
# nano-banana 1 does not support imageSize
|
||||
if "nano-banana-" in config.model:
|
||||
data["generationConfig"]["imageConfig"]["imageSize"] = config.resolution
|
||||
|
||||
model = nano_banana_models[config.model]
|
||||
url = f"{self.endpoint}/models/{model}:generateContent"
|
||||
|
||||
# Add system instruction if provided
|
||||
if len(system_message) > 0:
|
||||
data["systemInstruction"] = {"parts": [{"text": system_message[0].text}]}
|
||||
|
||||
response = requests.post(url, headers=headers, json=data, timeout=self.timeout)
|
||||
data: Dict = response.json()
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
raise Exception(data.get("error").get("message"))
|
||||
error_message = data.get("error").get("message", "Unknown error")
|
||||
return (f"Gemini API error: {error_message}", None)
|
||||
|
||||
return data["completion"]
|
||||
# Extract text and images from response
|
||||
if "candidates" in data and len(data["candidates"]) > 0:
|
||||
candidate = data["candidates"][0]
|
||||
if "content" in candidate and "parts" in candidate["content"]:
|
||||
parts = candidate["content"]["parts"]
|
||||
|
||||
# Collect text parts
|
||||
text_parts = []
|
||||
image_urls = []
|
||||
|
||||
for part in parts:
|
||||
# Handle text
|
||||
if "text" in part and part["text"]:
|
||||
text_parts.append(part["text"])
|
||||
|
||||
# Handle inline images
|
||||
elif "inlineData" in part and part["inlineData"]:
|
||||
image_urls.append(
|
||||
f"data:{part['inlineData']['mimeType']};base64,{part['inlineData']['data']}"
|
||||
)
|
||||
|
||||
text = (
|
||||
"".join(text_parts)
|
||||
if text_parts
|
||||
else "No text response from Gemini API"
|
||||
)
|
||||
images = image_urls_to_tensor(image_urls, timeout=self.timeout)
|
||||
|
||||
return (text, images)
|
||||
|
||||
return ("No response from Gemini API", None)
|
||||
|
||||
def complete(self, prompt: str, config: LLMConfig, seed=None):
|
||||
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
|
||||
return self.chat(messages, config, seed)
|
||||
|
||||
|
||||
class AwsBedrockMistralApi(BaseModel):
|
||||
@@ -215,7 +491,7 @@ class AwsBedrockMistralApi(BaseModel):
|
||||
def __init__(self, **data):
|
||||
super().__init__(**data)
|
||||
|
||||
ensure_package("boto3", version="1.34.101")
|
||||
ensure_package("boto3", required_version=">=1.34.101")
|
||||
import boto3
|
||||
|
||||
self.bedrock_runtime = boto3.client(
|
||||
@@ -240,13 +516,16 @@ class AwsBedrockMistralApi(BaseModel):
|
||||
"temperature": config.temperature,
|
||||
}
|
||||
|
||||
response = self.bedrock_runtime.invoke_model(body=json.dumps(data), modelId=config.model)
|
||||
response = self.bedrock_runtime.invoke_model(
|
||||
body=json.dumps(data), modelId=config.model
|
||||
)
|
||||
data: Dict = json.loads(response.get("body").read())
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
raise Exception(data.get("error").get("message"))
|
||||
return (f"Mistral API error: {data.get('error').get('message')}", None)
|
||||
|
||||
return data["outputs"][0]["text"]
|
||||
text = data["outputs"][0]["text"]
|
||||
return (text, None)
|
||||
|
||||
|
||||
class AwsBedrockClaudeApi(BaseModel):
|
||||
@@ -261,7 +540,7 @@ class AwsBedrockClaudeApi(BaseModel):
|
||||
def __init__(self, **data):
|
||||
super().__init__(**data)
|
||||
|
||||
ensure_package("boto3", version="1.34.101")
|
||||
ensure_package("boto3", required_version=">=1.34.101")
|
||||
import boto3
|
||||
|
||||
self.bedrock_runtime = boto3.client(
|
||||
@@ -273,7 +552,7 @@ class AwsBedrockClaudeApi(BaseModel):
|
||||
)
|
||||
|
||||
def chat(self, messages: List[LLMMessage], config: LLMConfig, seed=None):
|
||||
if config.model not in bedrock_claude3_models:
|
||||
if config.model not in bedrock_claude_models:
|
||||
raise Exception(f"Must provide a Claude v3 model, got {config.model}")
|
||||
|
||||
system_message = [m for m in messages if m.role == "system"]
|
||||
@@ -288,35 +567,23 @@ class AwsBedrockClaudeApi(BaseModel):
|
||||
"system": system_message[0].text if len(system_message) > 0 else None,
|
||||
}
|
||||
|
||||
response = self.bedrock_runtime.invoke_model(body=json.dumps(data), modelId=config.model)
|
||||
response = self.bedrock_runtime.invoke_model(
|
||||
body=json.dumps(data), modelId=config.model
|
||||
)
|
||||
data: Dict = json.loads(response.get("body").read())
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
raise Exception(data.get("error").get("message"))
|
||||
return (f"Claude API error: {data.get('error').get('message')}", None)
|
||||
|
||||
return data["content"][0]["text"]
|
||||
text = data["content"][0]["text"]
|
||||
return (text, None)
|
||||
|
||||
def complete(self, prompt: str, config: LLMConfig, seed=None):
|
||||
if config.model not in bedrock_claude2_models:
|
||||
raise Exception(f"Must provide a Claude v2 model, got {config.model}")
|
||||
|
||||
prompt = f"\n\nHuman: {prompt}\n\nAssistant:"
|
||||
data = {
|
||||
"prompt": prompt,
|
||||
"max_tokens_to_sample": config.max_token,
|
||||
"temperature": config.temperature,
|
||||
}
|
||||
|
||||
response = self.bedrock_runtime.invoke_model(body=json.dumps(data), modelId=config.model)
|
||||
data: Dict = json.loads(response.get("body").read())
|
||||
|
||||
if data.get("error", None) is not None:
|
||||
raise Exception(data.get("error").get("message"))
|
||||
|
||||
return data["completion"]
|
||||
messages = [LLMMessage(role=LLMMessageRole.user, text=prompt)]
|
||||
return self.chat(messages, config, seed)
|
||||
|
||||
|
||||
LLMApi = Union[OpenAIApi, ClaudeApi]
|
||||
LLMApi = Union[OpenAIApi, OpenRouterApi, ClaudeApi, GeminiApi]
|
||||
|
||||
|
||||
class OpenAIApiNode:
|
||||
@@ -325,7 +592,10 @@ class OpenAIApiNode:
|
||||
return {
|
||||
"required": {
|
||||
"openai_api_key": ("STRING", {"multiline": False}),
|
||||
"endpoint": ("STRING", {"multiline": False, "default": "https://api.openai.com/v1"}),
|
||||
"endpoint": (
|
||||
"STRING",
|
||||
{"multiline": False, "default": "https://api.openai.com/v1"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -342,13 +612,42 @@ class OpenAIApiNode:
|
||||
return (OpenAIApi(api_key=openai_api_key, endpoint=endpoint),)
|
||||
|
||||
|
||||
class OpenRouterApiNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"openrouter_api_key": ("STRING", {"multiline": False}),
|
||||
"endpoint": (
|
||||
"STRING",
|
||||
{"multiline": False, "default": "https://openrouter.ai/api/v1"},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LLM_API",)
|
||||
FUNCTION = "create_api"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def create_api(self, openrouter_api_key, endpoint):
|
||||
if not openrouter_api_key or openrouter_api_key == "":
|
||||
openrouter_api_key = os.environ.get("OPENROUTER_API_KEY")
|
||||
if not openrouter_api_key:
|
||||
raise Exception("OpenRouter API key is required.")
|
||||
|
||||
return (OpenRouterApi(api_key=openrouter_api_key, endpoint=endpoint),)
|
||||
|
||||
|
||||
class ClaudeApiNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"claude_api_key": ("STRING", {"multiline": False}),
|
||||
"endpoint": ("STRING", {"multiline": False, "default": "https://api.anthropic.com/v1"}),
|
||||
"endpoint": (
|
||||
"STRING",
|
||||
{"multiline": False, "default": "https://api.anthropic.com/v1"},
|
||||
),
|
||||
"version": (["2023-06-01"], {"default": "2023-06-01"}),
|
||||
},
|
||||
}
|
||||
@@ -360,13 +659,47 @@ class ClaudeApiNode:
|
||||
|
||||
def create_api(self, claude_api_key, endpoint, version):
|
||||
if not claude_api_key or claude_api_key == "":
|
||||
claude_api_key = os.environ.get("CLAUDE_API_KEY")
|
||||
claude_api_key = os.environ.get(
|
||||
"ANTHROPIC_API_KEY", os.environ.get("CLAUDE_API_KEY")
|
||||
)
|
||||
if not claude_api_key:
|
||||
raise Exception("Claude API key is required.")
|
||||
raise Exception("Anthropic API key is required.")
|
||||
|
||||
return (ClaudeApi(api_key=claude_api_key, endpoint=endpoint, version=version),)
|
||||
|
||||
|
||||
class GeminiApiNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"gemini_api_key": ("STRING", {"multiline": False}),
|
||||
"endpoint": (
|
||||
"STRING",
|
||||
{
|
||||
"multiline": False,
|
||||
"default": "https://generativelanguage.googleapis.com/v1beta",
|
||||
},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LLM_API",)
|
||||
RETURN_NAMES = ("llm_api",)
|
||||
FUNCTION = "create_api"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def create_api(self, gemini_api_key, endpoint):
|
||||
if not gemini_api_key or gemini_api_key == "":
|
||||
gemini_api_key = os.environ.get(
|
||||
"GEMINI_API_KEY", os.environ.get("GOOGLE_API_KEY")
|
||||
)
|
||||
if not gemini_api_key:
|
||||
raise Exception("Gemini API key is required.")
|
||||
|
||||
return (GeminiApi(api_key=gemini_api_key, endpoint=endpoint),)
|
||||
|
||||
|
||||
class AwsBedrockMistralApiNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
@@ -384,7 +717,9 @@ class AwsBedrockMistralApiNode:
|
||||
FUNCTION = "create_api"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def create_api(self, aws_access_key_id, aws_secret_access_key, aws_session_token, region):
|
||||
def create_api(
|
||||
self, aws_access_key_id, aws_secret_access_key, aws_session_token, region
|
||||
):
|
||||
if not aws_access_key_id or aws_access_key_id == "":
|
||||
aws_access_key_id = os.environ.get("AWS_ACCESS_KEY_ID", None)
|
||||
if not aws_secret_access_key or aws_secret_access_key == "":
|
||||
@@ -414,7 +749,10 @@ class AwsBedrockClaudeApiNode:
|
||||
"aws_secret_access_key": ("STRING", {"multiline": False}),
|
||||
"aws_session_token": ("STRING", {"multiline": False}),
|
||||
"region": (aws_regions, {"default": aws_regions[0]}),
|
||||
"version": (bedrock_anthropic_versions, {"default": bedrock_anthropic_versions[0]}),
|
||||
"version": (
|
||||
bedrock_anthropic_versions,
|
||||
{"default": bedrock_anthropic_versions[0]},
|
||||
),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -423,7 +761,14 @@ class AwsBedrockClaudeApiNode:
|
||||
FUNCTION = "create_api"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def create_api(self, aws_access_key_id, aws_secret_access_key, aws_session_token, region, version):
|
||||
def create_api(
|
||||
self,
|
||||
aws_access_key_id,
|
||||
aws_secret_access_key,
|
||||
aws_session_token,
|
||||
region,
|
||||
version,
|
||||
):
|
||||
if not aws_access_key_id or aws_access_key_id == "":
|
||||
aws_access_key_id = os.environ.get("AWS_ACCESS_KEY_ID", None)
|
||||
if not aws_secret_access_key or aws_secret_access_key == "":
|
||||
@@ -451,17 +796,18 @@ class LLMApiConfigNode:
|
||||
return {
|
||||
"required": {
|
||||
"model": (
|
||||
gpt_models
|
||||
+ claude3_models
|
||||
+ claude2_models
|
||||
+ bedrock_claude3_models
|
||||
+ bedrock_claude2_models
|
||||
+ bedrock_mistral_models,
|
||||
{"default": gpt_vision_models[0]},
|
||||
all_models,
|
||||
{"default": gpt_models[0]},
|
||||
),
|
||||
"max_token": ("INT", {"default": 1024}),
|
||||
"temperature": ("FLOAT", {"default": 0, "min": 0, "max": 1.0, "step": 0.001}),
|
||||
}
|
||||
"max_token": ("INT", {"default": 1024, "min": 1, "max": 102400}),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{"default": 0.5, "min": 0, "max": 1.0, "step": 0.001},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"custom_model": ("STRING", {"multiline": False, "default": ""})
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LLM_CONFIG",)
|
||||
@@ -469,8 +815,77 @@ class LLMApiConfigNode:
|
||||
FUNCTION = "make_config"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def make_config(self, max_token, model, temperature):
|
||||
return (LLMConfig(model=model, max_token=max_token, temperature=temperature),)
|
||||
def make_config(self, model, max_token, temperature, custom_model=""):
|
||||
# Use custom_model if provided, otherwise use the selected model from dropdown
|
||||
final_model = (
|
||||
custom_model.strip() if custom_model and custom_model.strip() else model
|
||||
)
|
||||
return (
|
||||
LLMConfig(model=final_model, max_token=max_token, temperature=temperature),
|
||||
)
|
||||
|
||||
|
||||
class NanoBananaApiConfigNode:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": (
|
||||
list(nano_banana_models.keys()),
|
||||
{"default": "nano-banana"},
|
||||
),
|
||||
"temperature": (
|
||||
"FLOAT",
|
||||
{"default": 0.5, "min": 0, "max": 1.0, "step": 0.001},
|
||||
),
|
||||
"modalities": (
|
||||
["image+text", "image", "text"],
|
||||
{"default": "image+text"},
|
||||
),
|
||||
"aspect_ratio": (
|
||||
[
|
||||
"auto",
|
||||
"1:1",
|
||||
"2:3",
|
||||
"3:2",
|
||||
"3:4",
|
||||
"4:3",
|
||||
"4:5",
|
||||
"5:4",
|
||||
"9:16",
|
||||
"16:9",
|
||||
"21:9",
|
||||
],
|
||||
{"default": "auto"},
|
||||
),
|
||||
"resolution": (["1K", "2K", "4K"], {"default": "1K"}),
|
||||
},
|
||||
"optional": {},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LLM_CONFIG",)
|
||||
RETURN_NAMES = ("llm_config",)
|
||||
FUNCTION = "make_config"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def make_config(
|
||||
self,
|
||||
model,
|
||||
temperature,
|
||||
modalities="image+text",
|
||||
aspect_ratio="auto",
|
||||
resolution="1K",
|
||||
):
|
||||
return (
|
||||
NanoBananaConfig(
|
||||
model=model,
|
||||
max_token=8192,
|
||||
temperature=temperature,
|
||||
modalities=modalities,
|
||||
aspect_ratio=aspect_ratio,
|
||||
resolution=resolution,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class LLMMessageNode:
|
||||
@@ -481,7 +896,13 @@ class LLMMessageNode:
|
||||
"role": (["system", "user", "assistant"],),
|
||||
"text": ("STRING", {"multiline": True}),
|
||||
},
|
||||
"optional": {"image": ("IMAGE",), "messages": ("LLM_MESSAGE",)},
|
||||
"optional": {
|
||||
"messages": ("LLM_MESSAGE",),
|
||||
"image": ("IMAGE",),
|
||||
"image_2": ("IMAGE",),
|
||||
"image_3": ("IMAGE",),
|
||||
"image_4": ("IMAGE",),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("LLM_MESSAGE",)
|
||||
@@ -489,21 +910,43 @@ class LLMMessageNode:
|
||||
FUNCTION = "make_message"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def make_message(self, role, text, image: Optional[Tensor] = None, messages: Optional[List[LLMMessage]] = None):
|
||||
def make_message(
|
||||
self,
|
||||
role,
|
||||
text,
|
||||
messages: Optional[List[LLMMessage]] = None,
|
||||
image: Optional[Tensor] = None,
|
||||
image_2: Optional[Tensor] = None,
|
||||
image_3: Optional[Tensor] = None,
|
||||
image_4: Optional[Tensor] = None,
|
||||
):
|
||||
messages = [] if messages is None else messages.copy()
|
||||
|
||||
if role == "system":
|
||||
if isinstance(image, Tensor):
|
||||
raise Exception("System prompt does not support image.")
|
||||
|
||||
system_message = [m for m in messages if m.role == "system"]
|
||||
if len(system_message) > 0:
|
||||
raise Exception("Only one system prompt is allowed.")
|
||||
|
||||
if isinstance(image, Tensor):
|
||||
pil = tensor2pil(image)
|
||||
content = pil2base64(pil)
|
||||
messages.append(LLMMessage(role=role, text=text, image=content))
|
||||
if any(
|
||||
isinstance(img, Tensor) for img in [image, image_2, image_3, image_4]
|
||||
):
|
||||
raise Exception("System prompt does not support image.")
|
||||
|
||||
all_images = []
|
||||
for img_tensor in [image, image_2, image_3, image_4]:
|
||||
if isinstance(img_tensor, Tensor):
|
||||
if len(img_tensor.shape) == 4: # Batch of images
|
||||
for i in range(img_tensor.shape[0]):
|
||||
pil = tensor2pil(img_tensor[i])
|
||||
content = pil2base64(pil)
|
||||
all_images.append(content)
|
||||
else: # Single image
|
||||
pil = tensor2pil(img_tensor)
|
||||
content = pil2base64(pil)
|
||||
all_images.append(content)
|
||||
|
||||
if all_images:
|
||||
messages.append(LLMMessage(role=role, text=text, images=all_images))
|
||||
else:
|
||||
messages.append(LLMMessage(role=role, text=text))
|
||||
|
||||
@@ -522,14 +965,14 @@ class LLMChatNode:
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("response",)
|
||||
RETURN_TYPES = ("STRING", "IMAGE")
|
||||
RETURN_NAMES = ("text", "(optional) images")
|
||||
FUNCTION = "chat"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def chat(self, messages: List[LLMMessage], api: LLMApi, config: LLMConfig, seed):
|
||||
response = api.chat(messages, config, seed)
|
||||
return (response,)
|
||||
text, images = api.chat(messages, config, seed)
|
||||
return (text, images)
|
||||
|
||||
|
||||
class LLMCompletionNode:
|
||||
@@ -544,22 +987,25 @@ class LLMCompletionNode:
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("response",)
|
||||
RETURN_TYPES = ("STRING", "IMAGE")
|
||||
RETURN_NAMES = ("text", "(optional) images")
|
||||
FUNCTION = "chat"
|
||||
CATEGORY = "ArtVenture/LLM"
|
||||
|
||||
def chat(self, prompt: str, api: LLMApi, config: LLMConfig, seed):
|
||||
response = api.complete(prompt, config, seed)
|
||||
return (response,)
|
||||
text, images = api.complete(prompt, config, seed)
|
||||
return (text, images)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"AV_OpenAIApi": OpenAIApiNode,
|
||||
"AV_OpenRouterApi": OpenRouterApiNode,
|
||||
"AV_ClaudeApi": ClaudeApiNode,
|
||||
"AV_GeminiApi": GeminiApiNode,
|
||||
"AV_AwsBedrockClaudeApi": AwsBedrockClaudeApiNode,
|
||||
"AV_AwsBedrockMistralApi": AwsBedrockMistralApiNode,
|
||||
"AV_LLMApiConfig": LLMApiConfigNode,
|
||||
"AV_NanoBananaApiConfig": NanoBananaApiConfigNode,
|
||||
"AV_LLMMessage": LLMMessageNode,
|
||||
"AV_LLMChat": LLMChatNode,
|
||||
"AV_LLMCompletion": LLMCompletionNode,
|
||||
@@ -567,10 +1013,13 @@ NODE_CLASS_MAPPINGS = {
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AV_OpenAIApi": "OpenAI API",
|
||||
"AV_OpenRouterApi": "OpenRouter API",
|
||||
"AV_ClaudeApi": "Claude API",
|
||||
"AV_GeminiApi": "Gemini API",
|
||||
"AV_AwsBedrockClaudeApi": "AWS Bedrock Claude API",
|
||||
"AV_AwsBedrockMistralApi": "AWS Bedrock Mistral API",
|
||||
"AV_LLMApiConfig": "LLM API Config",
|
||||
"AV_NanoBananaApiConfig": "NanoBanana API Config",
|
||||
"AV_LLMMessage": "LLM Message",
|
||||
"AV_LLMChat": "LLM Chat",
|
||||
"AV_LLMCompletion": "LLM Completion",
|
||||
|
||||
@@ -42,7 +42,7 @@ def load_file_from_url(
|
||||
*,
|
||||
model_dir: str,
|
||||
progress: bool = True,
|
||||
file_name: str | None = None,
|
||||
file_name: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Download a file from `url` into `model_dir`, using the file present if possible.
|
||||
|
||||
|
||||
+17
-24
@@ -70,7 +70,7 @@ class AVVAELoader(VAELoader):
|
||||
inputs["optional"] = {"vae_override": ("STRING", {"default": "None"})}
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def load_vae(self, vae_name, vae_override="None"):
|
||||
if vae_override != "None":
|
||||
@@ -92,7 +92,7 @@ class AVLoraLoader(LoraLoader):
|
||||
}
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def load_lora(self, model, clip, lora_name, *args, lora_override="None", enabled=True, **kwargs):
|
||||
if not enabled:
|
||||
@@ -119,7 +119,7 @@ class AVLoraListStacker:
|
||||
|
||||
RETURN_TYPES = ("LORA_STACK",)
|
||||
FUNCTION = "load_list_lora"
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
|
||||
def parse_lora_list(self, data: str):
|
||||
# data is a list of lora model (lora_name, strength_model, strength_clip, url) in json format
|
||||
@@ -213,7 +213,7 @@ class AVCheckpointModelsToParametersPipe:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PIPE",)
|
||||
CATEGORY = "Art Venture/Parameters"
|
||||
CATEGORY = "ArtVenture/Parameters"
|
||||
FUNCTION = "checkpoint_models_to_parameter_pipe"
|
||||
|
||||
def checkpoint_models_to_parameter_pipe(
|
||||
@@ -255,7 +255,7 @@ class AVPromptsToParametersPipe:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("PIPE",)
|
||||
CATEGORY = "Art Venture/Parameters"
|
||||
CATEGORY = "ArtVenture/Parameters"
|
||||
FUNCTION = "prompt_to_parameter_pipe"
|
||||
|
||||
def prompt_to_parameter_pipe(self, positive, negative, pipe: Dict = {}, image=None, mask=None):
|
||||
@@ -297,7 +297,7 @@ class AVParametersPipeToCheckpointModels:
|
||||
"lora_2_name",
|
||||
"lora_3_name",
|
||||
)
|
||||
CATEGORY = "Art Venture/Parameters"
|
||||
CATEGORY = "ArtVenture/Parameters"
|
||||
FUNCTION = "parameter_pipe_to_checkpoint_models"
|
||||
|
||||
def parameter_pipe_to_checkpoint_models(self, pipe: Dict = {}):
|
||||
@@ -346,7 +346,7 @@ class AVParametersPipeToPrompts:
|
||||
"image",
|
||||
"mask",
|
||||
)
|
||||
CATEGORY = "Art Venture/Parameters"
|
||||
CATEGORY = "ArtVenture/Parameters"
|
||||
FUNCTION = "parameter_pipe_to_prompt"
|
||||
|
||||
def parameter_pipe_to_prompt(self, pipe: Dict = {}):
|
||||
@@ -372,23 +372,23 @@ class AVCheckpointMerge:
|
||||
"model1": ("MODEL",),
|
||||
"model2": ("MODEL",),
|
||||
"model1_weight": ("FLOAT", {"default": 1.0, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
"model2_weight": ("FLOAT", {"default": 1.0, "min": -1.0, "max": 1.0, "step": 0.01}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
FUNCTION = "merge"
|
||||
|
||||
CATEGORY = "Art Venture/Model Merging"
|
||||
CATEGORY = "ArtVenture/Model Merging"
|
||||
DESCRIPTION = "DEPRECATED: Use ComfyUI's native ModelMergeSimple instead"
|
||||
|
||||
def merge(self, model1, model2, model1_weight, model2_weight):
|
||||
def merge(self, model1, model2, model1_weight):
|
||||
m = model1.clone()
|
||||
k1 = model1.get_key_patches("diffusion_model.")
|
||||
k2 = model2.get_key_patches("diffusion_model.")
|
||||
for k in k1:
|
||||
if k in k2:
|
||||
a = k1[k][0]
|
||||
b = k2[k][0]
|
||||
for k in k2:
|
||||
if k in k1:
|
||||
a, _ = k1[k][0]
|
||||
b, _ = k2[k][0]
|
||||
|
||||
if a.shape != b.shape and a.shape[0:1] + a.shape[2:] == b.shape[0:1] + b.shape[2:]:
|
||||
if a.shape[1] == 4 and b.shape[1] == 9:
|
||||
@@ -400,20 +400,13 @@ class AVCheckpointMerge:
|
||||
"When merging instruct-pix2pix model with a normal one, model1 must be the instruct-pix2pix model."
|
||||
)
|
||||
|
||||
c = torch.zeros_like(a)
|
||||
c[:, 0:4, :, :] = b
|
||||
b = c
|
||||
|
||||
m.add_patches({k: (b,)}, model2_weight, model1_weight)
|
||||
else:
|
||||
logger.warn(f"Key {k} not found in model2")
|
||||
m.add_patches({k: k1[k]}, -1.0, 1.0) # zero out
|
||||
m.add_patches({k: k2[k]}, 1 - model1_weight, model1_weight)
|
||||
|
||||
return (m,)
|
||||
|
||||
|
||||
class AVCheckpointSave(CheckpointSave):
|
||||
CATEGORY = "Art Venture/Model Merging"
|
||||
CATEGORY = "ArtVenture/Model Merging"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -462,7 +455,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"AV_LoraLoader": "Lora Loader",
|
||||
"AV_LoraListLoader": "Lora List Loader",
|
||||
"AV_LoraListStacker": "Lora List Stacker",
|
||||
"AV_CheckpointMerge": "Checkpoint Merge",
|
||||
"AV_CheckpointMerge": "[Deprecated] Checkpoint Merge",
|
||||
"AV_CheckpointSave": "Checkpoint Save",
|
||||
}
|
||||
|
||||
|
||||
@@ -37,17 +37,13 @@ class ColorBlend:
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "color_blending_mode"
|
||||
CATEGORY = "Art Venture/Post Processing"
|
||||
CATEGORY = "ArtVenture/Post Processing"
|
||||
|
||||
def color_blending_mode(self, bw_layer, color_layer):
|
||||
if bw_layer.shape[0] < color_layer.shape[0]:
|
||||
bw_layer = bw_layer.repeat(color_layer.shape[0], 1, 1, 1)[
|
||||
: color_layer.shape[0]
|
||||
]
|
||||
bw_layer = bw_layer.repeat(color_layer.shape[0], 1, 1, 1)[: color_layer.shape[0]]
|
||||
if bw_layer.shape[0] > color_layer.shape[0]:
|
||||
color_layer = color_layer.repeat(bw_layer.shape[0], 1, 1, 1)[
|
||||
: bw_layer.shape[0]
|
||||
]
|
||||
color_layer = color_layer.repeat(bw_layer.shape[0], 1, 1, 1)[: bw_layer.shape[0]]
|
||||
|
||||
batch_size, *_ = bw_layer.shape
|
||||
tensor_output = torch.empty_like(bw_layer)
|
||||
@@ -70,8 +66,6 @@ class ColorBlend:
|
||||
for i in range(batch_size):
|
||||
blend = color_blend(image1[i], image2[i])
|
||||
blend = np.stack([blend])
|
||||
tensor_output[i : i + 1] = (
|
||||
torch.from_numpy(blend.transpose(0, 3, 1, 2)) / 255.0
|
||||
).permute(0, 2, 3, 1)
|
||||
tensor_output[i : i + 1] = (torch.from_numpy(blend.transpose(0, 3, 1, 2)) / 255.0).permute(0, 2, 3, 1)
|
||||
|
||||
return (tensor_output,)
|
||||
|
||||
@@ -36,7 +36,7 @@ class ColorCorrect:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "color_correct"
|
||||
|
||||
CATEGORY = "Art Venture/Post Processing"
|
||||
CATEGORY = "ArtVenture/Post Processing"
|
||||
|
||||
def color_correct(
|
||||
self,
|
||||
|
||||
+134
-116
@@ -5,7 +5,7 @@ import torch
|
||||
import base64
|
||||
import random
|
||||
import requests
|
||||
from typing import List, Dict, Tuple
|
||||
from typing import List, Dict, Tuple, Optional
|
||||
|
||||
from PIL import Image, ImageOps, ImageFilter
|
||||
import numpy as np
|
||||
@@ -80,7 +80,7 @@ def prepare_image_for_preview(image: Image.Image, output_dir: str, prefix=None):
|
||||
|
||||
def load_images_from_url(urls: List[str], keep_alpha_channel=False):
|
||||
images: List[Image.Image] = []
|
||||
masks: List[Image.Image] = []
|
||||
masks: List[Optional[Image.Image]] = []
|
||||
|
||||
for url in urls:
|
||||
if url.startswith("data:image/"):
|
||||
@@ -158,11 +158,20 @@ class UtilLoadImageFromUrl:
|
||||
self.filename_prefix = "TempImageFromUrl"
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {},
|
||||
"required": {
|
||||
"image": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"placeholder": "Input image paths or URLS one per line. Eg:\nhttps://example.com/image.png\nfile:///path/to/local/image.jpg\ndata:image/png;base64,...",
|
||||
"multiline": True,
|
||||
"dynamicPrompts": False,
|
||||
},
|
||||
),
|
||||
},
|
||||
"optional": {
|
||||
"image": ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False}),
|
||||
"keep_alpha_channel": (
|
||||
"BOOLEAN",
|
||||
{"default": False, "label_on": "enabled", "label_off": "disabled"},
|
||||
@@ -171,76 +180,81 @@ class UtilLoadImageFromUrl:
|
||||
"BOOLEAN",
|
||||
{"default": False, "label_on": "list", "label_off": "batch"},
|
||||
),
|
||||
"url": ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "BOOLEAN")
|
||||
OUTPUT_IS_LIST = (True, True, False)
|
||||
RETURN_NAMES = ("images", "masks", "has_image")
|
||||
CATEGORY = "Art Venture/Image"
|
||||
CATEGORY = "ArtVenture/Image"
|
||||
FUNCTION = "load_image"
|
||||
|
||||
def load_image(self, image="", keep_alpha_channel=False, output_mode=False, url=""):
|
||||
if not image or image == "":
|
||||
image = url
|
||||
|
||||
def load_image(self, image: str, keep_alpha_channel=False, output_mode=False):
|
||||
urls = image.strip().split("\n")
|
||||
images, masks = load_images_from_url(urls, keep_alpha_channel)
|
||||
if len(images) == 0:
|
||||
image = torch.zeros((1, 64, 64, 3), dtype=torch.float32, device="cpu")
|
||||
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
images = [tensor2pil(image)]
|
||||
masks = [tensor2pil(mask, mode="L")]
|
||||
pil_images, pil_masks = load_images_from_url(urls, keep_alpha_channel)
|
||||
has_image = len(pil_images) > 0
|
||||
if not has_image:
|
||||
i = torch.zeros((1, 64, 64, 3), dtype=torch.float32, device="cpu")
|
||||
m = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
pil_images = [tensor2pil(i)]
|
||||
pil_masks = [tensor2pil(m, mode="L")]
|
||||
|
||||
previews = []
|
||||
np_images = []
|
||||
np_masks = []
|
||||
np_images: list[torch.Tensor] = []
|
||||
np_masks: list[torch.Tensor] = []
|
||||
|
||||
for image, mask in zip(images, masks):
|
||||
if mask is not None:
|
||||
preview_image = Image.new("RGB", image.size)
|
||||
preview_image.paste(image, (0, 0))
|
||||
preview_image.putalpha(mask)
|
||||
for pil_image, pil_mask in zip(pil_images, pil_masks):
|
||||
if pil_mask is not None:
|
||||
preview_image = Image.new("RGB", pil_image.size)
|
||||
preview_image.paste(pil_image, (0, 0))
|
||||
preview_image.putalpha(pil_mask)
|
||||
else:
|
||||
preview_image = image
|
||||
preview_image = pil_image
|
||||
|
||||
previews.append(prepare_image_for_preview(preview_image, self.output_dir, self.filename_prefix))
|
||||
|
||||
image = pil2tensor(image)
|
||||
if mask:
|
||||
mask = np.array(mask).astype(np.float32) / 255.0
|
||||
mask = 1.0 - torch.from_numpy(mask)
|
||||
np_image = pil2tensor(pil_image)
|
||||
if pil_mask:
|
||||
np_mask = np.array(pil_mask).astype(np.float32) / 255.0
|
||||
np_mask = 1.0 - torch.from_numpy(np_mask)
|
||||
else:
|
||||
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
np_mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
|
||||
np_images.append(image)
|
||||
np_masks.append(mask.unsqueeze(0))
|
||||
np_images.append(np_image)
|
||||
np_masks.append(np_mask.unsqueeze(0))
|
||||
|
||||
if output_mode:
|
||||
result = (np_images, np_masks, True)
|
||||
result = (np_images, np_masks, has_image)
|
||||
else:
|
||||
has_size_mismatch = False
|
||||
if len(np_images) > 1:
|
||||
for image in np_images[1:]:
|
||||
if image.shape[1] != np_images[0].shape[1] or image.shape[2] != np_images[0].shape[2]:
|
||||
for np_image in np_images[1:]:
|
||||
if np_image.shape[1] != np_images[0].shape[1] or np_image.shape[2] != np_images[0].shape[2]:
|
||||
has_size_mismatch = True
|
||||
break
|
||||
|
||||
if has_size_mismatch:
|
||||
raise Exception("To output as batch, images must have the same size. Use list output mode instead.")
|
||||
|
||||
result = ([torch.cat(np_images)], [torch.cat(np_masks)], True)
|
||||
result = ([torch.cat(np_images)], [torch.cat(np_masks)], has_image)
|
||||
|
||||
return {"ui": {"images": previews}, "result": result}
|
||||
|
||||
|
||||
class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image": ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False}),
|
||||
"image": (
|
||||
"STRING",
|
||||
{
|
||||
"default": "",
|
||||
"placeholder": "Input image paths or URLS one per line. Eg:\nhttps://example.com/image.png\nfile:///path/to/local/image.jpg\ndata:image/png;base64,...",
|
||||
"multiline": True,
|
||||
"dynamicPrompts": False,
|
||||
},
|
||||
),
|
||||
"channel": (["alpha", "red", "green", "blue"],),
|
||||
},
|
||||
"optional": {
|
||||
@@ -260,11 +274,11 @@ class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
|
||||
image = url
|
||||
|
||||
urls = image.strip().split("\n")
|
||||
images, alphas = load_images_from_url(urls, True)
|
||||
pil_images, pil_alphas = load_images_from_url(urls, True)
|
||||
|
||||
masks: List[torch.Tensor] = []
|
||||
|
||||
for img, alpha in zip(images, alphas):
|
||||
for img, alpha in zip(pil_images, pil_alphas):
|
||||
if channel == "alpha":
|
||||
mask = alpha
|
||||
elif channel == "red":
|
||||
@@ -297,7 +311,7 @@ class UtilLoadImageAsMaskFromUrl(UtilLoadImageFromUrl):
|
||||
|
||||
class UtilLoadJsonFromText:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"data": (
|
||||
@@ -308,7 +322,7 @@ class UtilLoadJsonFromText:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("JSON",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "load_json"
|
||||
|
||||
def load_json(self, data: str):
|
||||
@@ -317,7 +331,7 @@ class UtilLoadJsonFromText:
|
||||
|
||||
class UtilLoadJsonFromUrl:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"url": ("STRING", {"default": ""}),
|
||||
@@ -328,7 +342,7 @@ class UtilLoadJsonFromUrl:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("JSON",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "load_json"
|
||||
|
||||
def load_json(self, url: str, print_to_console=False):
|
||||
@@ -345,7 +359,7 @@ class UtilLoadJsonFromUrl:
|
||||
|
||||
class UtilGetObjectFromJson:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"json": ("JSON",),
|
||||
@@ -354,7 +368,7 @@ class UtilGetObjectFromJson:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("JSON",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_objects_from_json"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -364,7 +378,7 @@ class UtilGetObjectFromJson:
|
||||
|
||||
class UtilGetTextFromJson:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"json": ("JSON",),
|
||||
@@ -373,7 +387,7 @@ class UtilGetTextFromJson:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_string_from_json"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -383,7 +397,7 @@ class UtilGetTextFromJson:
|
||||
|
||||
class UtilGetFloatFromJson:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"json": ("JSON",),
|
||||
@@ -392,7 +406,7 @@ class UtilGetFloatFromJson:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_float_from_json"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -402,7 +416,7 @@ class UtilGetFloatFromJson:
|
||||
|
||||
class UtilGetIntFromJson:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"json": ("JSON",),
|
||||
@@ -411,7 +425,7 @@ class UtilGetIntFromJson:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_int_from_json"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -421,7 +435,7 @@ class UtilGetIntFromJson:
|
||||
|
||||
class UtilGetBoolFromJson:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"json": ("JSON",),
|
||||
@@ -430,7 +444,7 @@ class UtilGetBoolFromJson:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BOOLEAN",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_bool_from_json"
|
||||
OUTPUT_NODE = True
|
||||
|
||||
@@ -440,7 +454,7 @@ class UtilGetBoolFromJson:
|
||||
|
||||
class UtilRandomInt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"min": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
@@ -449,11 +463,11 @@ class UtilRandomInt:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", "STRING")
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "random_int"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, *args, **kwargs):
|
||||
def IS_CHANGED(cls, *args, **kwargs):
|
||||
return torch.rand(1).item()
|
||||
|
||||
def random_int(self, min: int, max: int):
|
||||
@@ -463,7 +477,7 @@ class UtilRandomInt:
|
||||
|
||||
class UtilRandomFloat:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"min": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
@@ -472,11 +486,11 @@ class UtilRandomFloat:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT", "STRING")
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "random_float"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, *args, **kwargs):
|
||||
def IS_CHANGED(cls, *args, **kwargs):
|
||||
return torch.rand(1).item()
|
||||
|
||||
def random_float(self, min: float, max: float):
|
||||
@@ -486,13 +500,13 @@ class UtilRandomFloat:
|
||||
|
||||
class UtilStringToInt:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {"string": ("STRING", {"default": "0"})},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "string_to_int"
|
||||
|
||||
def string_to_int(self, string: str):
|
||||
@@ -501,7 +515,7 @@ class UtilStringToInt:
|
||||
|
||||
class UtilStringToNumber:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"string": ("STRING", {"default": "0"}),
|
||||
@@ -510,7 +524,7 @@ class UtilStringToNumber:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("INT", "FLOAT")
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "string_to_numbers"
|
||||
|
||||
def string_to_numbers(self, string: str, rounding):
|
||||
@@ -526,7 +540,7 @@ class UtilStringToNumber:
|
||||
|
||||
class UtilNumberScaler:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"min": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 0xFFFFFFFFFFFFFFFF}),
|
||||
@@ -538,7 +552,7 @@ class UtilNumberScaler:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("FLOAT",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "scale_number"
|
||||
|
||||
def scale_number(self, min: float, max: float, scale_to_min: float, scale_to_max: float, value: float):
|
||||
@@ -548,7 +562,7 @@ class UtilNumberScaler:
|
||||
|
||||
class UtilBooleanPrimitive:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"value": ("BOOLEAN", {"default": False}),
|
||||
@@ -557,7 +571,7 @@ class UtilBooleanPrimitive:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BOOLEAN", "STRING")
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "boolean_primitive"
|
||||
|
||||
def boolean_primitive(self, value: bool, reverse: bool):
|
||||
@@ -569,7 +583,7 @@ class UtilBooleanPrimitive:
|
||||
|
||||
class UtilTextSwitchCase:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"switch_cases": (
|
||||
@@ -588,7 +602,7 @@ class UtilTextSwitchCase:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("STRING",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "text_switch_case"
|
||||
|
||||
def text_switch_case(self, switch_cases: str, condition: str, default_value: str, delimiter: str = ":"):
|
||||
@@ -602,7 +616,7 @@ class UtilTextSwitchCase:
|
||||
# Process previous case if exists
|
||||
if current_case is not None and condition == current_case:
|
||||
return ("\n".join(current_output),)
|
||||
|
||||
|
||||
# Start new case
|
||||
current_case, output = line.split(delimiter, 1)
|
||||
current_output = [output]
|
||||
@@ -618,7 +632,7 @@ class UtilTextSwitchCase:
|
||||
|
||||
class UtilImageMuxer:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image_1": ("IMAGE",),
|
||||
@@ -629,7 +643,7 @@ class UtilImageMuxer:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_muxer"
|
||||
|
||||
def image_muxer(self, image_1, image_2, input_selector, image_3=None, image_4=None):
|
||||
@@ -639,7 +653,7 @@ class UtilImageMuxer:
|
||||
|
||||
class UtilSDXLAspectRatioSelector:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"aspect_ratio": (
|
||||
@@ -665,7 +679,7 @@ class UtilSDXLAspectRatioSelector:
|
||||
RETURN_TYPES = ("STRING", "INT", "INT")
|
||||
RETURN_NAMES = ("ratio", "width", "height")
|
||||
FUNCTION = "get_aspect_ratio"
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
|
||||
def get_aspect_ratio(self, aspect_ratio):
|
||||
width, height = 1024, 1024
|
||||
@@ -702,7 +716,7 @@ class UtilSDXLAspectRatioSelector:
|
||||
|
||||
class UtilAspectRatioSelector(UtilSDXLAspectRatioSelector):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"aspect_ratio": (
|
||||
@@ -732,7 +746,7 @@ class UtilAspectRatioSelector(UtilSDXLAspectRatioSelector):
|
||||
|
||||
class UtilDependenciesEdit:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"dependencies": ("DEPENDENCIES",),
|
||||
@@ -758,7 +772,7 @@ class UtilDependenciesEdit:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("DEPENDENCIES",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "edit_dependencies"
|
||||
|
||||
def edit_dependencies(
|
||||
@@ -821,7 +835,7 @@ class UtilImageScaleDown:
|
||||
crop_methods = ["disabled", "center"]
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -833,12 +847,12 @@ class UtilImageScaleDown:
|
||||
"INT",
|
||||
{"default": 512, "min": 1, "max": MAX_RESOLUTION, "step": 1},
|
||||
),
|
||||
"crop": (s.crop_methods,),
|
||||
"crop": (cls.crop_methods,),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_scale_down"
|
||||
|
||||
def image_scale_down(self, images, width, height, crop):
|
||||
@@ -860,7 +874,7 @@ class UtilImageScaleDown:
|
||||
results = []
|
||||
for image in s:
|
||||
img = tensor2pil(image).convert("RGB")
|
||||
img = img.resize((width, height), Image.LANCZOS)
|
||||
img = img.resize((width, height), Image.Resampling.LANCZOS)
|
||||
results.append(pil2tensor(img))
|
||||
|
||||
return (torch.cat(results, dim=0),)
|
||||
@@ -868,7 +882,7 @@ class UtilImageScaleDown:
|
||||
|
||||
class UtilImageScaleDownBy(UtilImageScaleDown):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -880,7 +894,7 @@ class UtilImageScaleDownBy(UtilImageScaleDown):
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_scale_down_by"
|
||||
|
||||
def image_scale_down_by(self, images, scale_by):
|
||||
@@ -893,7 +907,7 @@ class UtilImageScaleDownBy(UtilImageScaleDown):
|
||||
|
||||
class UtilImageScaleDownToSize(UtilImageScaleDownBy):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -903,7 +917,7 @@ class UtilImageScaleDownToSize(UtilImageScaleDownBy):
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_scale_down_to_size"
|
||||
|
||||
def image_scale_down_to_size(self, images, size, mode):
|
||||
@@ -919,9 +933,13 @@ class UtilImageScaleDownToSize(UtilImageScaleDownBy):
|
||||
return self.image_scale_down_by(images, scale_by)
|
||||
|
||||
|
||||
class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
|
||||
class UtilImageScaleToTotalPixels(UtilImageScaleDownBy):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.upscale_model_node = ImageUpscaleWithModel()
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -933,7 +951,7 @@ class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_scale_down_to_total_pixels"
|
||||
|
||||
def image_scale_up_by(self, images: torch.Tensor, scale_by, upscale_model_opt):
|
||||
@@ -946,10 +964,10 @@ class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
|
||||
s = s.movedim(1, -1)
|
||||
return (s,)
|
||||
else:
|
||||
s = self.upscale(upscale_model_opt, images)[0]
|
||||
s = self.upscale_model_node.execute(upscale_model_opt, images)[0]
|
||||
return self.image_scale_down(s, width, height, "center")
|
||||
|
||||
def image_scale_down_to_total_pixels(self, images, megapixels, upscale_model_opt=None):
|
||||
def image_scale_down_to_total_pixels(self, images, megapixels, *args, upscale_model_opt=None, **kwargs):
|
||||
width = images.shape[2]
|
||||
height = images.shape[1]
|
||||
scale_by = np.sqrt((megapixels * 1024 * 1024) / (width * height))
|
||||
@@ -962,7 +980,7 @@ class UtilImageScaleToTotalPixels(UtilImageScaleDownBy, ImageUpscaleWithModel):
|
||||
|
||||
class UtilImageAlphaComposite:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"image_1": ("IMAGE",),
|
||||
@@ -971,7 +989,7 @@ class UtilImageAlphaComposite:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_alpha_composite"
|
||||
|
||||
def image_alpha_composite(self, image_1: torch.Tensor, image_2: torch.Tensor):
|
||||
@@ -994,7 +1012,7 @@ class UtilImageAlphaComposite:
|
||||
|
||||
class UtilImageGaussianBlur:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -1003,7 +1021,7 @@ class UtilImageGaussianBlur:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_gaussian_blur"
|
||||
|
||||
def image_gaussian_blur(self, images, radius):
|
||||
@@ -1018,7 +1036,7 @@ class UtilImageGaussianBlur:
|
||||
|
||||
class UtilImageExtractChannel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -1028,7 +1046,7 @@ class UtilImageExtractChannel:
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
RETURN_NAMES = ("channel_data",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_extract_alpha"
|
||||
|
||||
def image_extract_alpha(self, images: torch.Tensor, channel):
|
||||
@@ -1048,7 +1066,7 @@ class UtilImageExtractChannel:
|
||||
|
||||
class UtilImageApplyChannel:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -1058,7 +1076,7 @@ class UtilImageApplyChannel:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "image_apply_channel"
|
||||
|
||||
def image_apply_channel(self, images: torch.Tensor, channel_data: torch.Tensor, channel):
|
||||
@@ -1100,20 +1118,20 @@ class UtillQRCodeGenerator:
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
FUNCTION = "create_qr_code"
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
|
||||
def create_qr_code(self, text, size, qr_version, error_correction, box_size, border):
|
||||
ensure_package("qrcode", "qrcode[pil]")
|
||||
ensure_package("qrcode", install_package_name="qrcode[pil]")
|
||||
import qrcode
|
||||
|
||||
if error_correction == "L":
|
||||
error_level = qrcode.constants.ERROR_CORRECT_L
|
||||
error_level = qrcode.ERROR_CORRECT_L
|
||||
elif error_correction == "M":
|
||||
error_level = qrcode.constants.ERROR_CORRECT_M
|
||||
error_level = qrcode.ERROR_CORRECT_M
|
||||
elif error_correction == "Q":
|
||||
error_level = qrcode.constants.ERROR_CORRECT_Q
|
||||
error_level = qrcode.ERROR_CORRECT_Q
|
||||
else:
|
||||
error_level = qrcode.constants.ERROR_CORRECT_H
|
||||
error_level = qrcode.ERROR_CORRECT_H
|
||||
|
||||
qr = qrcode.QRCode(version=qr_version, error_correction=error_level, box_size=box_size, border=border)
|
||||
qr.add_data(text)
|
||||
@@ -1126,7 +1144,7 @@ class UtillQRCodeGenerator:
|
||||
|
||||
class UtilRepeatImages:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"images": ("IMAGE",),
|
||||
@@ -1135,7 +1153,7 @@ class UtilRepeatImages:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "rebatch"
|
||||
|
||||
def rebatch(self, images: torch.Tensor, amount):
|
||||
@@ -1144,7 +1162,7 @@ class UtilRepeatImages:
|
||||
|
||||
class UtilSeedSelector:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"mode": ("BOOLEAN", {"default": True, "label_on": "random", "label_off": "fixed"}),
|
||||
@@ -1158,7 +1176,7 @@ class UtilSeedSelector:
|
||||
|
||||
RETURN_TYPES = ("INT",)
|
||||
RETURN_NAMES = ("seed",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_seed"
|
||||
|
||||
def get_seed(self, mode, seed, fixed_seed):
|
||||
@@ -1167,7 +1185,7 @@ class UtilSeedSelector:
|
||||
|
||||
class UtilCheckpointSelector:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"ckpt_name": (folder_paths.get_filename_list("checkpoints"),),
|
||||
@@ -1176,11 +1194,11 @@ class UtilCheckpointSelector:
|
||||
|
||||
RETURN_TYPES = (folder_paths.get_filename_list("checkpoints"), "STRING")
|
||||
RETURN_NAMES = ("ckpt_name", "ckpt_name_str")
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "get_ckpt_name"
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, *args, **kwargs):
|
||||
def IS_CHANGED(cls, *args, **kwargs):
|
||||
return torch.rand(1).item()
|
||||
|
||||
def get_ckpt_name(self, ckpt_name):
|
||||
@@ -1189,7 +1207,7 @@ class UtilCheckpointSelector:
|
||||
|
||||
class UtilModelMerge:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model1": ("MODEL",),
|
||||
@@ -1199,7 +1217,7 @@ class UtilModelMerge:
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MODEL",)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "merge_models"
|
||||
|
||||
def merge_models(self, model1, model2, ratio=1.0):
|
||||
@@ -1221,7 +1239,7 @@ class UtilModelMerge:
|
||||
|
||||
class UtilTextRandomMultiline:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"text": ("STRING", {"multiline": True, "dynamicPrompts": False}),
|
||||
@@ -1233,7 +1251,7 @@ class UtilTextRandomMultiline:
|
||||
RETURN_TYPES = ("STRING",)
|
||||
RETURN_NAMES = ("lines",)
|
||||
OUTPUT_IS_LIST = (True,)
|
||||
CATEGORY = "Art Venture/Utils"
|
||||
CATEGORY = "ArtVenture/Utils"
|
||||
FUNCTION = "random_multiline"
|
||||
|
||||
def random_multiline(self, text: str, amount=1, seed=0):
|
||||
|
||||
+29
-11
@@ -5,9 +5,10 @@ import torch
|
||||
import base64
|
||||
import numpy as np
|
||||
import importlib
|
||||
import importlib.metadata
|
||||
import subprocess
|
||||
import pkg_resources
|
||||
from pkg_resources import parse_version
|
||||
from packaging import version
|
||||
from packaging.specifiers import SpecifierSet
|
||||
from PIL import Image
|
||||
|
||||
from .logger import logger
|
||||
@@ -21,23 +22,40 @@ class AnyType(str):
|
||||
any_type = AnyType("*")
|
||||
|
||||
|
||||
def ensure_package(package, version=None, install_package_name=None):
|
||||
def ensure_package(package, required_version=None, install_package_name=None):
|
||||
# Try to import the package
|
||||
try:
|
||||
module = importlib.import_module(package)
|
||||
except ImportError:
|
||||
logger.info(f"Package {package} is not installed. Installing now...")
|
||||
install_command = _construct_pip_command(install_package_name or package, version)
|
||||
install_command = _construct_pip_command(install_package_name or package, required_version)
|
||||
subprocess.check_call(install_command)
|
||||
else:
|
||||
# If a specific version is required, check the version
|
||||
if version:
|
||||
installed_version = pkg_resources.get_distribution(package).version
|
||||
if parse_version(installed_version) < parse_version(version):
|
||||
logger.info(
|
||||
f"Package {package} is outdated (installed: {installed_version}, required: {version}). Upgrading now..."
|
||||
)
|
||||
install_command = _construct_pip_command(install_package_name or package, version)
|
||||
if required_version:
|
||||
try:
|
||||
installed_version = importlib.metadata.version(package)
|
||||
|
||||
# Parse version specifier (e.g., ">=1.1.1", "==1.1.1", "<=1.1.1")
|
||||
if any(op in required_version for op in ['>=', '<=', '==', '!=', '>', '<', '~=']):
|
||||
spec = SpecifierSet(required_version)
|
||||
if installed_version not in spec:
|
||||
logger.info(
|
||||
f"Package {package} version constraint not satisfied (installed: {installed_version}, required: {required_version}). Installing now..."
|
||||
)
|
||||
install_command = _construct_pip_command(install_package_name or package, required_version)
|
||||
subprocess.check_call(install_command)
|
||||
else:
|
||||
# Fallback to simple version comparison for backwards compatibility
|
||||
if version.parse(installed_version) < version.parse(required_version):
|
||||
logger.info(
|
||||
f"Package {package} is outdated (installed: {installed_version}, required: {required_version}). Upgrading now..."
|
||||
)
|
||||
install_command = _construct_pip_command(install_package_name or package, required_version)
|
||||
subprocess.check_call(install_command)
|
||||
except importlib.metadata.PackageNotFoundError:
|
||||
logger.info(f"Package {package} version information not found. Installing required version {required_version}...")
|
||||
install_command = _construct_pip_command(install_package_name or package, required_version)
|
||||
subprocess.check_call(install_command)
|
||||
|
||||
|
||||
|
||||
@@ -52,10 +52,10 @@ try:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
inputs = LoadVideoPath.INPUT_TYPES()
|
||||
inputs["required"]["video"] = ("STRING", {"default": "", "multiline": True, "dynamicPrompts": False})
|
||||
inputs["required"]["video"] = ("STRING", {"default": ""})
|
||||
return inputs
|
||||
|
||||
CATEGORY = "Art Venture/Loaders"
|
||||
CATEGORY = "ArtVenture/Loaders"
|
||||
FUNCTION = "load"
|
||||
RETURN_TYPES = ("IMAGE", "INT", "BOOLEAN")
|
||||
RETURN_NAMES = ("frames", "frame_count", "has_video")
|
||||
@@ -164,7 +164,7 @@ try:
|
||||
from urllib.parse import parse_qs
|
||||
|
||||
qs_idx = url.find("?")
|
||||
qs = parse_qs(url[qs_idx + 1:])
|
||||
qs = parse_qs(url[qs_idx + 1 :])
|
||||
filename = qs.get("name", qs.get("filename", None))
|
||||
if filename is None:
|
||||
raise Exception(f"Invalid url: {url}")
|
||||
|
||||
+14
-2
@@ -1,9 +1,21 @@
|
||||
[project]
|
||||
name = "comfyui-art-venture"
|
||||
description = "A comprehensive set of custom nodes for ComfyUI, focusing on utilities for image processing, JSON manipulation, model operations and working with object via URLs"
|
||||
version = "1.0.6"
|
||||
version = "1.1.7"
|
||||
license = "LICENSE"
|
||||
dependencies = ["timm==0.6.13", "transformers", "fairscale", "pycocoevalcap", "opencv-python", "qrcode[pil]", "pytorch_lightning", "kornia", "pydantic", "segment_anything", "boto3>=1.34.101"]
|
||||
dependencies = [
|
||||
"timm==0.6.13",
|
||||
"transformers",
|
||||
"fairscale",
|
||||
"pycocoevalcap",
|
||||
"opencv-python",
|
||||
"qrcode[pil]",
|
||||
"pytorch_lightning",
|
||||
"kornia",
|
||||
"pydantic",
|
||||
"segment_anything",
|
||||
"boto3>=1.34.101",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Repository = "https://github.com/sipherxyz/comfyui-art-venture"
|
||||
|
||||
+1
-1
@@ -9,4 +9,4 @@ kornia
|
||||
pydantic
|
||||
segment_anything
|
||||
omegaconf
|
||||
boto3>=1.34.101
|
||||
boto3
|
||||
|
||||
@@ -1,10 +1,9 @@
|
||||
import { app } from '../../../scripts/app.js';
|
||||
import { ComfyWidgets } from '../../../scripts/widgets.js';
|
||||
import { app } from '../../scripts/app.js';
|
||||
import { ComfyWidgets } from '../../scripts/widgets.js';
|
||||
|
||||
import {
|
||||
addKVState,
|
||||
chainCallback,
|
||||
hideWidgetForGood,
|
||||
addWidgetChangeCallback,
|
||||
} from './utils.js';
|
||||
|
||||
|
||||
+212
-521
@@ -1,547 +1,244 @@
|
||||
import { app, ANIM_PREVIEW_WIDGET } from '../../../scripts/app.js';
|
||||
import { api } from '../../../scripts/api.js';
|
||||
import { $el } from '../../../scripts/ui.js';
|
||||
import { createImageHost } from '../../../scripts/ui/imagePreview.js';
|
||||
import { app } from '../../scripts/app.js';
|
||||
import { api } from '../../scripts/api.js';
|
||||
import { $el } from '../../scripts/ui.js';
|
||||
import { addWidget, DOMWidgetImpl } from '../../scripts/domWidget.js';
|
||||
import { ComfyWidgets } from '../../scripts/widgets.js'
|
||||
|
||||
import { chainCallback, addKVState } from './utils.js';
|
||||
import { chainCallback, addKVState, addWidgetChangeCallback } from './utils.js';
|
||||
|
||||
const style = `
|
||||
.comfy-img-preview video {
|
||||
object-fit: contain;
|
||||
width: var(--comfy-img-preview-width);
|
||||
height: var(--comfy-img-preview-height);
|
||||
}
|
||||
`;
|
||||
const supportedNodes = ['LoadImageFromUrl', 'LoadImageAsMaskFromUrl'];
|
||||
|
||||
const URL_REGEX = /^((blob:)?https?:\/\/|\/view\?|\/api\/view\?|data:image\/)/
|
||||
const formatUrl = (url) => {
|
||||
if (!url) return ""
|
||||
|
||||
const supportedNodes = ['LoadImageFromUrl', 'LoadImageAsMaskFromUrl', 'LoadVideoFromUrl'];
|
||||
if (url.startsWith("http://") || url.startsWith("https://") || url.startsWith("blob:")) return url
|
||||
if (url.startsWith("/view") || url.startsWith("/api/view")) return url
|
||||
|
||||
function injectHidden(widget) {
|
||||
widget.computeSize = (target_width) => {
|
||||
if (widget.hidden) {
|
||||
return [0, -4];
|
||||
}
|
||||
return [target_width, 20];
|
||||
};
|
||||
widget._type = widget.type;
|
||||
Object.defineProperty(widget, 'type', {
|
||||
set: function (value) {
|
||||
widget._type = value;
|
||||
},
|
||||
get: function () {
|
||||
if (widget.hidden) {
|
||||
return 'hidden';
|
||||
}
|
||||
return widget._type;
|
||||
},
|
||||
});
|
||||
}
|
||||
|
||||
function migrateWidget(nodeType, oldWidgetName, newWidgetName) {
|
||||
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
|
||||
if (!this.widgets) return;
|
||||
|
||||
const oldIndex = this.widgets.findIndex((w) => w.name === oldWidgetName);
|
||||
if (oldIndex > -1) {
|
||||
this.widgets.splice(oldIndex, 1);
|
||||
}
|
||||
|
||||
chainCallback(this, 'onConfigure', function (info) {
|
||||
if (typeof info.widgets_values != 'object') return;
|
||||
|
||||
const newWidget = this.widgets.find((w) => w.name === newWidgetName);
|
||||
if (newWidget && info.widgets_values[oldWidgetName]) {
|
||||
newWidget.value = info.widgets_values[oldWidgetName];
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
|
||||
function formatImageUrl(params) {
|
||||
if (typeof params === "string") {
|
||||
if (URL_REGEX.test(params)) return params;
|
||||
|
||||
const folder_separator = params.lastIndexOf("/");
|
||||
let subfolder = "";
|
||||
if (folder_separator > -1) {
|
||||
subfolder = params.substring(0, folder_separator);
|
||||
params = params.substring(folder_separator + 1);
|
||||
}
|
||||
let type = "input";
|
||||
if (params.indexOf(" [") > -1) {
|
||||
type = params.split(" [")[1].split("]")[0];
|
||||
params = params.split(" [")[0];
|
||||
}
|
||||
|
||||
params = {
|
||||
filename: params,
|
||||
type: type,
|
||||
subfolder: subfolder,
|
||||
};
|
||||
let type = "input"
|
||||
if (url.endsWith(']')) {
|
||||
const openBracketIndex = url.lastIndexOf('[')
|
||||
type = url.slice(openBracketIndex + 1, url.length - 1).trim()
|
||||
url = url.slice(0, openBracketIndex).trim()
|
||||
}
|
||||
|
||||
if (params.url) {
|
||||
return params.url;
|
||||
}
|
||||
const parts = url.split('/')
|
||||
const filename = parts.pop()
|
||||
const subfolder = parts.join('/')
|
||||
|
||||
params = { ...params };
|
||||
const params = [
|
||||
'filename=' + encodeURIComponent(filename),
|
||||
'type=' + type,
|
||||
'subfolder=' + subfolder,
|
||||
app.getRandParam().substring(1)
|
||||
].join('&')
|
||||
|
||||
if (!params.filename && params.name) {
|
||||
params.filename = params.name;
|
||||
delete params.name;
|
||||
}
|
||||
|
||||
return api.apiURL("/view?" + new URLSearchParams(params).toString() + app.getPreviewFormatParam());
|
||||
return api.apiURL(`/view?${params}`)
|
||||
}
|
||||
|
||||
async function uploadFile(file) {
|
||||
//TODO: Add uploaded file to cache with Cache.put()?
|
||||
try {
|
||||
// Wrap file in formdata so it includes filename
|
||||
const body = new FormData();
|
||||
const i = file.webkitRelativePath.lastIndexOf('/');
|
||||
const subfolder = file.webkitRelativePath.slice(0, i + 1);
|
||||
const new_file = new File([file], file.name, {
|
||||
type: file.type,
|
||||
lastModified: file.lastModified,
|
||||
});
|
||||
body.append('image', new_file);
|
||||
if (i > 0) {
|
||||
body.append('subfolder', subfolder);
|
||||
}
|
||||
const resp = await api.fetchApi('/upload/image', {
|
||||
method: 'POST',
|
||||
body,
|
||||
});
|
||||
// copied from ComfyUI_frontend/src/composables/widgets/useStringWidget.ts
|
||||
// remove the Object.defineProperty(widget, 'value') part
|
||||
function addUrlWidget(node, name, options) {
|
||||
const inputEl = document.createElement('textarea')
|
||||
inputEl.className = 'comfy-multiline-input'
|
||||
inputEl.value = options.default
|
||||
inputEl.placeholder = options.placeholder || name
|
||||
inputEl.spellcheck = false
|
||||
|
||||
if (resp.status === 200 || resp.status === 201) {
|
||||
return resp.json();
|
||||
} else {
|
||||
alert(`Upload failed: ${resp.statusText}`);
|
||||
}
|
||||
} catch (error) {
|
||||
alert(`Upload failed: ${error}`);
|
||||
}
|
||||
}
|
||||
|
||||
function addVideoCustomSize(nodeType, nodeData, widgetName) {
|
||||
//Add the extra size widgets now
|
||||
//This takes some finagling as widget order is defined by key order
|
||||
const newWidgets = {};
|
||||
for (let key in nodeData.input.required) {
|
||||
newWidgets[key] = nodeData.input.required[key];
|
||||
if (key == widgetName) {
|
||||
newWidgets[key][0] = newWidgets[key][0].concat(['Custom Width', 'Custom Height', 'Custom']);
|
||||
newWidgets['custom_width'] = ['INT', { default: 512, min: 8, step: 8 }];
|
||||
newWidgets['custom_height'] = ['INT', { default: 512, min: 8, step: 8 }];
|
||||
}
|
||||
}
|
||||
nodeData.input.required = newWidgets;
|
||||
|
||||
//Add a callback which sets up the actual logic once the node is created
|
||||
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
|
||||
const node = this;
|
||||
const sizeOptionWidget = node.widgets.find((w) => w.name === widgetName);
|
||||
const widthWidget = node.widgets.find((w) => w.name === 'custom_width');
|
||||
const heightWidget = node.widgets.find((w) => w.name === 'custom_height');
|
||||
injectHidden(widthWidget);
|
||||
widthWidget.options.serialize = false;
|
||||
injectHidden(heightWidget);
|
||||
heightWidget.options.serialize = false;
|
||||
sizeOptionWidget._value = sizeOptionWidget.value;
|
||||
Object.defineProperty(sizeOptionWidget, 'value', {
|
||||
set: function (value) {
|
||||
//TODO: Only modify hidden/reset size when a change occurs
|
||||
if (value == 'Custom Width') {
|
||||
widthWidget.hidden = false;
|
||||
heightWidget.hidden = true;
|
||||
} else if (value == 'Custom Height') {
|
||||
widthWidget.hidden = true;
|
||||
heightWidget.hidden = false;
|
||||
} else if (value == 'Custom') {
|
||||
widthWidget.hidden = false;
|
||||
heightWidget.hidden = false;
|
||||
} else {
|
||||
widthWidget.hidden = true;
|
||||
heightWidget.hidden = true;
|
||||
}
|
||||
node.setSize([node.size[0], node.computeSize([node.size[0], node.size[1]])[1]]);
|
||||
this._value = value;
|
||||
const widget = new DOMWidgetImpl({
|
||||
node,
|
||||
name,
|
||||
type: 'customtext',
|
||||
element: inputEl,
|
||||
options: {
|
||||
hideOnZoom: true,
|
||||
getValue() {
|
||||
return inputEl.value
|
||||
},
|
||||
get: function () {
|
||||
return this._value;
|
||||
},
|
||||
});
|
||||
//Ensure proper visibility/size state for initial value
|
||||
sizeOptionWidget.value = sizeOptionWidget._value;
|
||||
setValue(v) {
|
||||
inputEl.value = v
|
||||
}
|
||||
}
|
||||
})
|
||||
addWidget(node, widget)
|
||||
|
||||
sizeOptionWidget.serializeValue = function () {
|
||||
if (this.value == 'Custom Width') {
|
||||
return widthWidget.value + 'x?';
|
||||
} else if (this.value == 'Custom Height') {
|
||||
return '?x' + heightWidget;
|
||||
} else if (this.value == 'Custom') {
|
||||
return widthWidget.value + 'x' + heightWidget.value;
|
||||
widget.inputEl = inputEl
|
||||
widget.options.minNodeSize = [400, 200]
|
||||
|
||||
inputEl.addEventListener('input', () => {
|
||||
widget.value = inputEl.value
|
||||
widget.callback?.(inputEl.value, true)
|
||||
})
|
||||
|
||||
// Allow middle mouse button panning
|
||||
inputEl.addEventListener('pointerdown', (event) => {
|
||||
if (event.button === 1) {
|
||||
app.canvas.processMouseDown(event)
|
||||
}
|
||||
})
|
||||
|
||||
inputEl.addEventListener('pointermove', (event) => {
|
||||
if ((event.buttons & 4) === 4) {
|
||||
app.canvas.processMouseMove(event)
|
||||
}
|
||||
})
|
||||
|
||||
inputEl.addEventListener('pointerup', (event) => {
|
||||
if (event.button === 1) {
|
||||
app.canvas.processMouseUp(event)
|
||||
}
|
||||
})
|
||||
|
||||
/** Timer reference. `null` when the timer completes. */
|
||||
let ignoreEventsTimer = null
|
||||
/** Total number of events ignored since the timer started. */
|
||||
let ignoredEvents = 0
|
||||
|
||||
// Pass wheel events to the canvas when appropriate
|
||||
inputEl.addEventListener('wheel', (event) => {
|
||||
if (!Object.is(event.deltaX, -0)) return
|
||||
|
||||
// If the textarea has focus, require more effort to activate pass-through
|
||||
const multiplier = document.activeElement === inputEl ? 2 : 1
|
||||
const maxScrollHeight = inputEl.scrollHeight - inputEl.clientHeight
|
||||
|
||||
if (
|
||||
(event.deltaY < 0 && inputEl.scrollTop === 0) ||
|
||||
(event.deltaY > 0 && inputEl.scrollTop === maxScrollHeight)
|
||||
) {
|
||||
// Attempting to scroll past the end of the textarea
|
||||
if (!ignoreEventsTimer || ignoredEvents > 25 * multiplier) {
|
||||
app.canvas.processMouseWheel(event)
|
||||
} else {
|
||||
return this.value;
|
||||
ignoredEvents++
|
||||
}
|
||||
};
|
||||
});
|
||||
} else if (event.deltaY !== 0) {
|
||||
// Start timer whenever a successful scroll occurs
|
||||
ignoredEvents = 0
|
||||
if (ignoreEventsTimer) clearTimeout(ignoreEventsTimer)
|
||||
|
||||
ignoreEventsTimer = setTimeout(() => {
|
||||
ignoreEventsTimer = null
|
||||
}, 800 * multiplier)
|
||||
}
|
||||
})
|
||||
|
||||
return widget
|
||||
}
|
||||
|
||||
function addUploadWidget(nodeType, widgetName, type) {
|
||||
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
|
||||
this.images = [];
|
||||
const pathWidget = this.widgets.find((w) => w.name === widgetName);
|
||||
const supportMultiple = pathWidget.type === 'customtext';
|
||||
function addImageUploadWidget(nodeType, nodeData, imageInputName) {
|
||||
const { input } = nodeData ?? {}
|
||||
const required = input?.required
|
||||
if (!required) return
|
||||
|
||||
if (pathWidget.element) {
|
||||
pathWidget.options.getMinHeight = () => 50;
|
||||
pathWidget.options.getMaxHeight = () => 150;
|
||||
}
|
||||
const imageOptions = required.image
|
||||
delete required.image
|
||||
|
||||
const fileInput = document.createElement('input');
|
||||
chainCallback(this, 'onRemoved', () => {
|
||||
fileInput?.remove();
|
||||
});
|
||||
chainCallback(nodeType.prototype, "onNodeCreated", function () {
|
||||
this.previewMediaType = 'image'
|
||||
|
||||
if (type === 'image') {
|
||||
Object.assign(fileInput, {
|
||||
type: 'file',
|
||||
accept: 'image/png,image/jpeg,image/webp',
|
||||
style: 'display: none',
|
||||
multiple: supportMultiple,
|
||||
onchange: async () => {
|
||||
if (!fileInput.files.length) {
|
||||
return;
|
||||
}
|
||||
const urlWidget = addUrlWidget(this, imageInputName, imageOptions[1])
|
||||
// move urlWidget to the first position
|
||||
const widgets = this.widgets.filter(w => w !== urlWidget)
|
||||
this.widgets = [urlWidget, ...widgets]
|
||||
|
||||
let successes = [];
|
||||
for (const file of fileInput.files) {
|
||||
const params = await uploadFile(file);
|
||||
ComfyWidgets.IMAGEUPLOAD(
|
||||
this,
|
||||
'upload',
|
||||
["IMAGEUPLOAD", { "image_upload": true, imageInputName }],
|
||||
)
|
||||
|
||||
if (!!params) {
|
||||
successes.push(params);
|
||||
} else {
|
||||
// Upload failed, but some prior uploads may have succeeded
|
||||
// Stop future uploads to prevent cascading failures
|
||||
// and only add to list if an upload has succeeded
|
||||
if (successes.length) {
|
||||
break;
|
||||
} else {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pathWidget.value = successes.map(formatImageUrl).join('\n');
|
||||
fileInput.value = '';
|
||||
},
|
||||
});
|
||||
} else if (type === 'video') {
|
||||
Object.assign(fileInput, {
|
||||
type: 'file',
|
||||
accept: 'video/webm,video/mp4,video/mkv,image/gif,image/webp',
|
||||
style: 'display: none',
|
||||
multiple: supportMultiple,
|
||||
onchange: async () => {
|
||||
if (!fileInput.files.length) {
|
||||
return;
|
||||
}
|
||||
|
||||
let successes = [];
|
||||
for (const file of fileInput.files) {
|
||||
const params = await uploadFile(file);
|
||||
|
||||
if (!!params) {
|
||||
successes.push(params);
|
||||
} else {
|
||||
// Upload failed, but some prior uploads may have succeeded
|
||||
// Stop future uploads to prevent cascading failures
|
||||
// and only add to list if an upload has succeeded
|
||||
if (successes.length) {
|
||||
break;
|
||||
} else {
|
||||
return;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pathWidget.value = successes.map(formatImageUrl).join('\n');
|
||||
fileInput.value = '';
|
||||
},
|
||||
});
|
||||
} else {
|
||||
throw new Error(`Unknown upload type ${type}`);
|
||||
}
|
||||
|
||||
document.body.append(fileInput);
|
||||
let uploadWidget = this.addWidget('button', 'choose ' + type + ' to upload', 'image', () => {
|
||||
//clear the active click event
|
||||
app.canvas.node_widget = null;
|
||||
fileInput.click();
|
||||
});
|
||||
uploadWidget.serialize = false;
|
||||
|
||||
// Add handler to check if an image is being dragged over our node
|
||||
this.onDragOver = function (e) {
|
||||
if (e.dataTransfer && e.dataTransfer.items) {
|
||||
const image = [...e.dataTransfer.items].find((f) => f.kind === 'file');
|
||||
return !!image;
|
||||
}
|
||||
|
||||
return false;
|
||||
};
|
||||
|
||||
// On drop upload files
|
||||
this.onDragDrop = async function (e) {
|
||||
let successes = [];
|
||||
const files = e.dataTransfer.files
|
||||
.filter((file) => file.type.startsWith('image/'))
|
||||
.slice(0, supportMultiple ? undefined : 1);
|
||||
|
||||
for (const file of files) {
|
||||
const params = await uploadFile(file);
|
||||
if (!!params) {
|
||||
successes.push(params);
|
||||
}
|
||||
pathWidget.value = (supportMultiple ? this.images : [])
|
||||
.concat(...successes.map(formatImageUrl))
|
||||
.join('\n');
|
||||
}
|
||||
|
||||
return successes.length > 0;
|
||||
};
|
||||
|
||||
this.pasteFile = function (file) {
|
||||
if (file.type.startsWith('image/')) {
|
||||
uploadFile(file).then((res) => {
|
||||
pathWidget.value = (supportMultiple ? this.images : [])
|
||||
.concat(formatImageUrl(res))
|
||||
.join('\n');
|
||||
});
|
||||
return true;
|
||||
}
|
||||
return false;
|
||||
};
|
||||
});
|
||||
}
|
||||
|
||||
function patchValueSetter(nodeType, widgetName) {
|
||||
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
|
||||
const pathWidget = this.widgets.find((w) => w.name === widgetName);
|
||||
pathWidget._value = pathWidget.value;
|
||||
let editing = false;
|
||||
|
||||
const setter = (value) => {
|
||||
if (typeof value !== 'string') value = formatImageUrl(value);
|
||||
|
||||
pathWidget._value = value;
|
||||
this.images = (value ?? '').split('\n').filter(Boolean);
|
||||
if (pathWidget.type === 'customtext' && !editing) {
|
||||
pathWidget.inputEl.value = value;
|
||||
}
|
||||
delete app.nodeOutputs[this.id]
|
||||
};
|
||||
|
||||
Object.defineProperty(pathWidget, 'value', {
|
||||
set: setter,
|
||||
get: () => pathWidget._value,
|
||||
});
|
||||
|
||||
if (pathWidget.type === 'customtext') {
|
||||
pathWidget.inputEl.addEventListener('focus', (e) => {
|
||||
editing = true;
|
||||
});
|
||||
pathWidget.inputEl.addEventListener('blur', (e) => {
|
||||
editing = false;
|
||||
});
|
||||
pathWidget.inputEl.addEventListener('keyup', (e) => {
|
||||
setter(e.target.value);
|
||||
const safeLoadImageFromUrl = (url) => {
|
||||
return new Promise((resolve, reject) => {
|
||||
const img = new Image();
|
||||
img.onload = () => resolve(img);
|
||||
img.onerror = () => reject(null);
|
||||
img.src = url;
|
||||
});
|
||||
}
|
||||
|
||||
pathWidget.callback = setter;
|
||||
pathWidget.value = pathWidget._value;
|
||||
});
|
||||
}
|
||||
let initialImgs = undefined;
|
||||
let isUrlImageSet = false
|
||||
|
||||
function addVideoPreview(nodeType, widgetName) {
|
||||
const createVideoNode = (url) => {
|
||||
return new Promise((cb) => {
|
||||
const videoEl = document.createElement('video');
|
||||
Object.defineProperty(videoEl, 'naturalWidth', {
|
||||
get: () => {
|
||||
return videoEl.videoWidth;
|
||||
},
|
||||
});
|
||||
Object.defineProperty(videoEl, 'naturalHeight', {
|
||||
get: () => {
|
||||
return videoEl.videoHeight;
|
||||
},
|
||||
});
|
||||
videoEl.addEventListener('loadedmetadata', () => {
|
||||
videoEl.controls = false;
|
||||
videoEl.loop = true;
|
||||
videoEl.muted = true;
|
||||
cb(videoEl);
|
||||
});
|
||||
videoEl.addEventListener('error', () => {
|
||||
cb();
|
||||
});
|
||||
videoEl.src = url;
|
||||
});
|
||||
};
|
||||
|
||||
const createImageNode = (url) => {
|
||||
return new Promise((cb) => {
|
||||
const imgEl = document.createElement('img');
|
||||
imgEl.onload = () => {
|
||||
cb(imgEl);
|
||||
};
|
||||
imgEl.addEventListener('error', () => {
|
||||
cb();
|
||||
});
|
||||
imgEl.src = url;
|
||||
});
|
||||
};
|
||||
|
||||
nodeType.prototype.onDrawBackground = function (ctx) {
|
||||
if (this.flags.collapsed) return;
|
||||
|
||||
let imageURLs = (this.images ?? []).map(formatImageUrl);
|
||||
let imagesChanged = false;
|
||||
|
||||
if (JSON.stringify(this.displayingImages) !== JSON.stringify(imageURLs)) {
|
||||
this.displayingImages = imageURLs;
|
||||
imagesChanged = true;
|
||||
}
|
||||
|
||||
if (!imagesChanged) return;
|
||||
if (!imageURLs.length) {
|
||||
this.imgs = null;
|
||||
this.animatedImages = false;
|
||||
return;
|
||||
}
|
||||
|
||||
const promises = imageURLs.map((url) => {
|
||||
if (/^(\/api)?\/view/.test(url)) {
|
||||
url = window.location.origin + url;
|
||||
}
|
||||
|
||||
let ext = '';
|
||||
if (url.startsWith('data:')) {
|
||||
const blob = dataUriToBlob(url);
|
||||
ext = blob.type.split('/').pop();
|
||||
url = URL.createObjectURL(blob);
|
||||
} else {
|
||||
const u = new URL(url);
|
||||
const filename =
|
||||
u.searchParams.get('filename') ||
|
||||
u.searchParams.get('name') ||
|
||||
u.pathname.split('/').pop();
|
||||
ext = filename.split('.').pop();
|
||||
}
|
||||
|
||||
const format = ['gif', 'webp', 'avif'].includes(ext) ? 'image' : 'video';
|
||||
if (format === 'video') {
|
||||
return createVideoNode(url);
|
||||
} else {
|
||||
return createImageNode(url);
|
||||
}
|
||||
});
|
||||
|
||||
Promise.all(promises)
|
||||
.then((imgs) => {
|
||||
this.imgs = imgs.filter(Boolean);
|
||||
})
|
||||
.then(() => {
|
||||
if (!this.imgs.length) return;
|
||||
|
||||
this.animatedImages = true;
|
||||
const widgetIdx = this.widgets?.findIndex((w) => w.name === ANIM_PREVIEW_WIDGET);
|
||||
|
||||
// Instead of using the canvas we'll use a IMG
|
||||
if (widgetIdx > -1) {
|
||||
// Replace content
|
||||
const widget = this.widgets[widgetIdx];
|
||||
widget.options.host.updateImages(this.imgs);
|
||||
} else {
|
||||
const host = createImageHost(this);
|
||||
this.setSizeForImage(true);
|
||||
const widget = this.addDOMWidget(ANIM_PREVIEW_WIDGET, 'img', host.el, {
|
||||
host,
|
||||
getHeight: host.getHeight,
|
||||
onDraw: host.onDraw,
|
||||
hideOnZoom: false,
|
||||
});
|
||||
widget.serializeValue = () => ({
|
||||
height: host.el.clientHeight,
|
||||
});
|
||||
// widget.computeSize = (w) => ([w, 220]);
|
||||
|
||||
widget.options.host.updateImages(this.imgs);
|
||||
}
|
||||
|
||||
this.imgs.forEach((img) => {
|
||||
if (img instanceof HTMLVideoElement) {
|
||||
img.muted = true;
|
||||
img.autoplay = true;
|
||||
img.play();
|
||||
}
|
||||
});
|
||||
});
|
||||
};
|
||||
|
||||
patchValueSetter(nodeType, widgetName);
|
||||
|
||||
chainCallback(nodeType.prototype, 'onExecuted', function (message) {
|
||||
if (message?.videos) {
|
||||
this.images = message?.videos.map(formatImageUrl);
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function addImagePreview(nodeType, widgetName) {
|
||||
const onDrawBackground = nodeType.prototype.onDrawBackground;
|
||||
nodeType.prototype.onDrawBackground = function (ctx) {
|
||||
if (this.flags.collapsed) return;
|
||||
|
||||
let imageURLs = (this.images ?? []).map(formatImageUrl);
|
||||
let imagesChanged = false;
|
||||
|
||||
if (JSON.stringify(this.displayingImages) !== JSON.stringify(imageURLs)) {
|
||||
this.displayingImages = imageURLs;
|
||||
imagesChanged = true;
|
||||
}
|
||||
|
||||
if (imagesChanged) {
|
||||
const setImagesFromUrl = (value = "") => {
|
||||
this.imageIndex = null;
|
||||
if (imageURLs.length > 0) {
|
||||
Promise.all(
|
||||
imageURLs.map((src) => {
|
||||
return new Promise((r) => {
|
||||
const img = new Image();
|
||||
img.onload = () => r(img);
|
||||
img.onerror = () => r(null);
|
||||
img.src = src;
|
||||
});
|
||||
}),
|
||||
).then((imgs) => {
|
||||
this.imgs = imgs.filter(Boolean);
|
||||
this.setSizeForImage?.();
|
||||
app.graph.setDirtyCanvas(true);
|
||||
});
|
||||
} else {
|
||||
this.imgs = null;
|
||||
|
||||
const urls = value.split("\n").filter(Boolean).map(formatUrl);
|
||||
if (!urls.length) {
|
||||
this.imgs = undefined;
|
||||
this.widgets = this.widgets.filter((w) => w.name !== "$$canvas-image-preview");
|
||||
isUrlImageSet = true;
|
||||
return
|
||||
}
|
||||
|
||||
return Promise.all(
|
||||
urls.map(safeLoadImageFromUrl)
|
||||
).then((imgs) => {
|
||||
initialImgs = imgs.filter(Boolean);
|
||||
this.imgs = initialImgs.length > 0 ? initialImgs : undefined;
|
||||
if (!this.imgs) {
|
||||
this.widgets = this.widgets.filter((w) => w.name !== "$$canvas-image-preview");
|
||||
}
|
||||
app.graph.setDirtyCanvas(true);
|
||||
|
||||
// cancel any img change in the next 2 seconds
|
||||
// to prevent `ComfyUI `overwriting the image
|
||||
setTimeout(() => {
|
||||
isUrlImageSet = true;
|
||||
}, 2000);
|
||||
|
||||
return initialImgs;
|
||||
})
|
||||
}
|
||||
|
||||
onDrawBackground?.call(this, ctx);
|
||||
};
|
||||
addWidgetChangeCallback(this, {
|
||||
name: "imgs",
|
||||
shouldChange: (value) => {
|
||||
if (isUrlImageSet) return true;
|
||||
return value === initialImgs;
|
||||
}
|
||||
})
|
||||
|
||||
patchValueSetter(nodeType, widgetName);
|
||||
urlWidget.callback = (value) => {
|
||||
if (!value) { // from upload
|
||||
value = urlWidget.value.split("\n").filter(Boolean).map(formatUrl).join("\n")
|
||||
}
|
||||
if (Array.isArray(value)) {
|
||||
value = value.map(formatUrl).join("\n")
|
||||
}
|
||||
if (value !== urlWidget.options.getValue()) {
|
||||
urlWidget.options.setValue(value)
|
||||
}
|
||||
setImagesFromUrl(value)
|
||||
}
|
||||
|
||||
this.clipspace = () => {
|
||||
const widgets = this.widgets
|
||||
.map(({ type, name, value }) => ({
|
||||
type,
|
||||
name,
|
||||
value,
|
||||
}))
|
||||
|
||||
widgets.push(
|
||||
{
|
||||
type: "text",
|
||||
name: imageInputName,
|
||||
value: (urlWidget.value || "").split("\n").filter(Boolean)[0],
|
||||
},
|
||||
{
|
||||
type: "text",
|
||||
name: "url",
|
||||
value: (urlWidget.value || "").split("\n").filter(Boolean)[0],
|
||||
}
|
||||
)
|
||||
|
||||
return { widgets, images: undefined }
|
||||
}
|
||||
|
||||
requestAnimationFrame(() => {
|
||||
setImagesFromUrl(urlWidget.value);
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
@@ -558,16 +255,10 @@ app.registerExtension({
|
||||
return;
|
||||
}
|
||||
|
||||
addKVState(nodeType);
|
||||
|
||||
if (nodeData.name === 'LoadImageFromUrl' || nodeData.name === 'LoadImageAsMaskFromUrl') {
|
||||
migrateWidget(nodeType, 'url', 'image');
|
||||
addUploadWidget(nodeType, 'image', 'image');
|
||||
addImagePreview(nodeType, 'image');
|
||||
} else if (nodeData.name == 'LoadVideoFromUrl') {
|
||||
addVideoCustomSize(nodeType, nodeData, 'force_size');
|
||||
addUploadWidget(nodeType, 'video', 'video');
|
||||
addVideoPreview(nodeType, 'video');
|
||||
addImageUploadWidget(nodeType, nodeData, 'image');
|
||||
}
|
||||
|
||||
addKVState(nodeType);
|
||||
},
|
||||
});
|
||||
|
||||
+63
-63
@@ -1,146 +1,146 @@
|
||||
export const CONVERTED_TYPE = "converted-widget";
|
||||
export const CONVERTED_TYPE = "converted-widget"
|
||||
|
||||
export function hideWidgetForGood(node, widget, suffix = "") {
|
||||
widget.origType = widget.type;
|
||||
widget.origComputeSize = widget.computeSize;
|
||||
widget.computeSize = () => [0, -4]; // -4 is due to the gap litegraph adds between widgets automatically
|
||||
widget.type = CONVERTED_TYPE + suffix;
|
||||
widget.origType = widget.type
|
||||
widget.origComputeSize = widget.computeSize
|
||||
widget.computeSize = () => [0, -4] // -4 is due to the gap litegraph adds between widgets automatically
|
||||
widget.type = CONVERTED_TYPE + suffix
|
||||
|
||||
// Hide any linked widgets, e.g. seed+seedControl
|
||||
if (widget.linkedWidgets) {
|
||||
for (const w of widget.linkedWidgets) {
|
||||
hideWidgetForGood(node, w, ":" + widget.name);
|
||||
hideWidgetForGood(node, w, ":" + widget.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
const doesInputWithNameExist = (node, name) => {
|
||||
return node.inputs ? node.inputs.some((input) => input.name === name) : false;
|
||||
};
|
||||
return node.inputs ? node.inputs.some(input => input.name === name) : false
|
||||
}
|
||||
|
||||
const HIDDEN_TAG = "tschide";
|
||||
const origProps = {};
|
||||
const HIDDEN_TAG = "tschide"
|
||||
const origProps = {}
|
||||
|
||||
// Toggle Widget + change size
|
||||
export function toggleWidget(node, widget, show = false, suffix = "", updateSize = true) {
|
||||
if (!widget || doesInputWithNameExist(node, widget.name)) return;
|
||||
if (!widget || doesInputWithNameExist(node, widget.name)) return
|
||||
|
||||
// Store the original properties of the widget if not already stored
|
||||
if (!origProps[widget.name]) {
|
||||
origProps[widget.name] = {
|
||||
origType: widget.type,
|
||||
origComputeSize: widget.computeSize,
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
const origSize = node.size;
|
||||
const origSize = node.size
|
||||
|
||||
// Set the widget type and computeSize based on the show flag
|
||||
widget.type = show ? origProps[widget.name].origType : HIDDEN_TAG + suffix;
|
||||
widget.computeSize = show
|
||||
? origProps[widget.name].origComputeSize
|
||||
: () => [0, -4];
|
||||
widget.type = show ? origProps[widget.name].origType : HIDDEN_TAG + suffix
|
||||
widget.computeSize = show ? origProps[widget.name].origComputeSize : () => [0, -4]
|
||||
|
||||
// Recursively handle linked widgets if they exist
|
||||
widget.linkedWidgets?.forEach((w) =>
|
||||
toggleWidget(node, w, ":" + widget.name, show)
|
||||
);
|
||||
widget.linkedWidgets?.forEach(w => toggleWidget(node, w, ":" + widget.name, show))
|
||||
|
||||
// Calculate the new height for the node based on its computeSize method
|
||||
if (updateSize) {
|
||||
const newHeight = node.computeSize()[1];
|
||||
node.setSize([node.size[0], newHeight]);
|
||||
const newHeight = node.computeSize()[1]
|
||||
node.setSize([node.size[0], newHeight])
|
||||
}
|
||||
}
|
||||
|
||||
export function addWidgetChangeCallback(widget, callback) {
|
||||
let widgetValue = widget.value;
|
||||
let originalDescriptor = Object.getOwnPropertyDescriptor(widget, "value");
|
||||
let widgetValue = widget.value
|
||||
let originalDescriptor = Object.getOwnPropertyDescriptor(widget, "value")
|
||||
Object.defineProperty(widget, "value", {
|
||||
get() {
|
||||
return originalDescriptor && originalDescriptor.get
|
||||
? originalDescriptor.get.call(widget)
|
||||
: widgetValue;
|
||||
return originalDescriptor && originalDescriptor.get ? originalDescriptor.get.call(widget) : widgetValue
|
||||
},
|
||||
set(newVal) {
|
||||
if (originalDescriptor && originalDescriptor.set) {
|
||||
originalDescriptor.set.call(widget, newVal);
|
||||
originalDescriptor.set.call(widget, newVal)
|
||||
} else {
|
||||
widgetValue = newVal;
|
||||
widgetValue = newVal
|
||||
}
|
||||
|
||||
callback(newVal);
|
||||
callback(newVal)
|
||||
},
|
||||
});
|
||||
})
|
||||
}
|
||||
|
||||
export function chainCallback(object, property, callback) {
|
||||
if (object == undefined) {
|
||||
//This should not happen.
|
||||
console.error("Tried to add callback to non-existant object");
|
||||
return;
|
||||
console.error("Tried to add callback to non-existant object")
|
||||
return
|
||||
}
|
||||
if (property in object) {
|
||||
const callback_orig = object[property];
|
||||
const callback_orig = object[property]
|
||||
object[property] = function () {
|
||||
const r = callback_orig.apply(this, arguments);
|
||||
callback.apply(this, arguments);
|
||||
return r;
|
||||
};
|
||||
const r = callback_orig?.apply(this, arguments)
|
||||
callback.apply(this, arguments)
|
||||
return r
|
||||
}
|
||||
} else {
|
||||
object[property] = callback;
|
||||
object[property] = callback
|
||||
}
|
||||
}
|
||||
|
||||
export function addKVState(nodeType) {
|
||||
chainCallback(nodeType.prototype, 'onNodeCreated', function () {
|
||||
chainCallback(this, 'onConfigure', function (info) {
|
||||
chainCallback(nodeType.prototype, "onNodeCreated", function () {
|
||||
chainCallback(this, "onConfigure", function (info) {
|
||||
if (!this.widgets) {
|
||||
//Node has no widgets, there is nothing to restore
|
||||
return;
|
||||
return
|
||||
}
|
||||
if (typeof info.widgets_values != 'object') {
|
||||
if (typeof info.widgets_values != "object") {
|
||||
//widgets_values is in some unknown inactionable format
|
||||
return;
|
||||
return
|
||||
}
|
||||
let widgetDict = info.widgets_values;
|
||||
let widgetDict = info.widgets_values
|
||||
if (widgetDict.length == undefined) {
|
||||
for (let w of this.widgets) {
|
||||
if (w.name in widgetDict) {
|
||||
w.value = widgetDict[w.name];
|
||||
w.value = widgetDict[w.name]
|
||||
if (w.type !== "button") {
|
||||
w.callback?.(w.value)
|
||||
}
|
||||
} else {
|
||||
//attempt to restore default value
|
||||
let inputs = LiteGraph.getNodeType(this.type).nodeData.input;
|
||||
let initialValue = null;
|
||||
let inputs = LiteGraph.getNodeType(this.type).nodeData.input
|
||||
let initialValue = null
|
||||
if (inputs?.required?.hasOwnProperty(w.name)) {
|
||||
if (inputs.required[w.name][1]?.hasOwnProperty('default')) {
|
||||
initialValue = inputs.required[w.name][1].default;
|
||||
if (inputs.required[w.name][1]?.hasOwnProperty("default")) {
|
||||
initialValue = inputs.required[w.name][1].default
|
||||
} else if (inputs.required[w.name][0].length) {
|
||||
initialValue = inputs.required[w.name][0][0];
|
||||
initialValue = inputs.required[w.name][0][0]
|
||||
}
|
||||
} else if (inputs?.optional?.hasOwnProperty(w.name)) {
|
||||
if (inputs.optional[w.name][1]?.hasOwnProperty('default')) {
|
||||
initialValue = inputs.optional[w.name][1].default;
|
||||
if (inputs.optional[w.name][1]?.hasOwnProperty("default")) {
|
||||
initialValue = inputs.optional[w.name][1].default
|
||||
} else if (inputs.optional[w.name][0].length) {
|
||||
initialValue = inputs.optional[w.name][0][0];
|
||||
initialValue = inputs.optional[w.name][0][0]
|
||||
}
|
||||
}
|
||||
if (initialValue) {
|
||||
w.value = initialValue;
|
||||
w.value = initialValue
|
||||
if (w.type !== "button") {
|
||||
w.callback?.(w.value)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
chainCallback(this, 'onSerialize', function (info) {
|
||||
info.widgets_values = {};
|
||||
})
|
||||
chainCallback(this, "onSerialize", function (info) {
|
||||
info.widgets_values = {}
|
||||
if (!this.widgets) {
|
||||
//object has no widgets, there is nothing to store
|
||||
return;
|
||||
return
|
||||
}
|
||||
for (let w of this.widgets) {
|
||||
info.widgets_values[w.name] = w.value;
|
||||
info.widgets_values[w.name] = w.value
|
||||
}
|
||||
});
|
||||
});
|
||||
}
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user