mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-13 07:47:06 +08:00
Merge branch 'master' into upstream-model-manager
This commit is contained in:
commit
06ce62de4b
@ -3,8 +3,8 @@ name: Python Linting
|
|||||||
on: [push, pull_request]
|
on: [push, pull_request]
|
||||||
|
|
||||||
jobs:
|
jobs:
|
||||||
pylint:
|
ruff:
|
||||||
name: Run Pylint
|
name: Run Ruff
|
||||||
runs-on: ubuntu-latest
|
runs-on: ubuntu-latest
|
||||||
|
|
||||||
steps:
|
steps:
|
||||||
@ -16,8 +16,8 @@ jobs:
|
|||||||
with:
|
with:
|
||||||
python-version: 3.x
|
python-version: 3.x
|
||||||
|
|
||||||
- name: Install Pylint
|
- name: Install Ruff
|
||||||
run: pip install pylint
|
run: pip install ruff
|
||||||
|
|
||||||
- name: Run Pylint
|
- name: Run Ruff
|
||||||
run: pylint --rcfile=.pylintrc $(find . -type f -name "*.py")
|
run: ruff check .
|
||||||
53
.github/workflows/test-ci.yml
vendored
53
.github/workflows/test-ci.yml
vendored
@ -20,7 +20,8 @@ jobs:
|
|||||||
strategy:
|
strategy:
|
||||||
fail-fast: false
|
fail-fast: false
|
||||||
matrix:
|
matrix:
|
||||||
os: [macos, linux, windows]
|
# os: [macos, linux, windows]
|
||||||
|
os: [macos, linux]
|
||||||
python_version: ["3.9", "3.10", "3.11", "3.12"]
|
python_version: ["3.9", "3.10", "3.11", "3.12"]
|
||||||
cuda_version: ["12.1"]
|
cuda_version: ["12.1"]
|
||||||
torch_version: ["stable"]
|
torch_version: ["stable"]
|
||||||
@ -31,9 +32,9 @@ jobs:
|
|||||||
- os: linux
|
- os: linux
|
||||||
runner_label: [self-hosted, Linux]
|
runner_label: [self-hosted, Linux]
|
||||||
flags: ""
|
flags: ""
|
||||||
- os: windows
|
# - os: windows
|
||||||
runner_label: [self-hosted, Windows]
|
# runner_label: [self-hosted, Windows]
|
||||||
flags: ""
|
# flags: ""
|
||||||
runs-on: ${{ matrix.runner_label }}
|
runs-on: ${{ matrix.runner_label }}
|
||||||
steps:
|
steps:
|
||||||
- name: Test Workflows
|
- name: Test Workflows
|
||||||
@ -45,28 +46,28 @@ jobs:
|
|||||||
google_credentials: ${{ secrets.GCS_SERVICE_ACCOUNT_JSON }}
|
google_credentials: ${{ secrets.GCS_SERVICE_ACCOUNT_JSON }}
|
||||||
comfyui_flags: ${{ matrix.flags }}
|
comfyui_flags: ${{ matrix.flags }}
|
||||||
|
|
||||||
test-win-nightly:
|
# test-win-nightly:
|
||||||
strategy:
|
# strategy:
|
||||||
fail-fast: true
|
# fail-fast: true
|
||||||
matrix:
|
# matrix:
|
||||||
os: [windows]
|
# os: [windows]
|
||||||
python_version: ["3.9", "3.10", "3.11", "3.12"]
|
# python_version: ["3.9", "3.10", "3.11", "3.12"]
|
||||||
cuda_version: ["12.1"]
|
# cuda_version: ["12.1"]
|
||||||
torch_version: ["nightly"]
|
# torch_version: ["nightly"]
|
||||||
include:
|
# include:
|
||||||
- os: windows
|
# - os: windows
|
||||||
runner_label: [self-hosted, Windows]
|
# runner_label: [self-hosted, Windows]
|
||||||
flags: ""
|
# flags: ""
|
||||||
runs-on: ${{ matrix.runner_label }}
|
# runs-on: ${{ matrix.runner_label }}
|
||||||
steps:
|
# steps:
|
||||||
- name: Test Workflows
|
# - name: Test Workflows
|
||||||
uses: comfy-org/comfy-action@main
|
# uses: comfy-org/comfy-action@main
|
||||||
with:
|
# with:
|
||||||
os: ${{ matrix.os }}
|
# os: ${{ matrix.os }}
|
||||||
python_version: ${{ matrix.python_version }}
|
# python_version: ${{ matrix.python_version }}
|
||||||
torch_version: ${{ matrix.torch_version }}
|
# torch_version: ${{ matrix.torch_version }}
|
||||||
google_credentials: ${{ secrets.GCS_SERVICE_ACCOUNT_JSON }}
|
# google_credentials: ${{ secrets.GCS_SERVICE_ACCOUNT_JSON }}
|
||||||
comfyui_flags: ${{ matrix.flags }}
|
# comfyui_flags: ${{ matrix.flags }}
|
||||||
|
|
||||||
test-unix-nightly:
|
test-unix-nightly:
|
||||||
strategy:
|
strategy:
|
||||||
|
|||||||
2
.github/workflows/test-launch.yml
vendored
2
.github/workflows/test-launch.yml
vendored
@ -28,7 +28,7 @@ jobs:
|
|||||||
- name: Start ComfyUI server
|
- name: Start ComfyUI server
|
||||||
run: |
|
run: |
|
||||||
python main.py --cpu 2>&1 | tee console_output.log &
|
python main.py --cpu 2>&1 | tee console_output.log &
|
||||||
wait-for-it --service 127.0.0.1:8188 -t 600
|
wait-for-it --service 127.0.0.1:8188 -t 30
|
||||||
working-directory: ComfyUI
|
working-directory: ComfyUI
|
||||||
- name: Check for unhandled exceptions in server log
|
- name: Check for unhandled exceptions in server log
|
||||||
run: |
|
run: |
|
||||||
|
|||||||
23
CODEOWNERS
23
CODEOWNERS
@ -1 +1,22 @@
|
|||||||
* @comfyanonymous
|
# Admins
|
||||||
|
* @comfyanonymous
|
||||||
|
|
||||||
|
# Note: Github teams syntax cannot be used here as the repo is not owned by Comfy-Org.
|
||||||
|
# Inlined the team members for now.
|
||||||
|
|
||||||
|
# Maintainers
|
||||||
|
*.md @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata @Kosinkadink
|
||||||
|
/tests/ @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata @Kosinkadink
|
||||||
|
/tests-unit/ @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata @Kosinkadink
|
||||||
|
/notebooks/ @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata @Kosinkadink
|
||||||
|
/script_examples/ @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata @Kosinkadink
|
||||||
|
|
||||||
|
# Python web server
|
||||||
|
/api_server/ @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata
|
||||||
|
/app/ @yoland68 @robinjhuang @huchenlei @webfiltered @pythongosssss @ltdrdata
|
||||||
|
|
||||||
|
# Frontend assets
|
||||||
|
/web/ @huchenlei @webfiltered @pythongosssss
|
||||||
|
|
||||||
|
# Extra nodes
|
||||||
|
/comfy_extras/ @yoland68 @robinjhuang @huchenlei @pythongosssss @ltdrdata @Kosinkadink
|
||||||
|
|||||||
@ -213,6 +213,14 @@ For 6700, 6600 and maybe other RDNA2 or older: ```HSA_OVERRIDE_GFX_VERSION=10.3.
|
|||||||
|
|
||||||
For AMD 7600 and maybe other RDNA3 cards: ```HSA_OVERRIDE_GFX_VERSION=11.0.0 python main.py```
|
For AMD 7600 and maybe other RDNA3 cards: ```HSA_OVERRIDE_GFX_VERSION=11.0.0 python main.py```
|
||||||
|
|
||||||
|
### AMD ROCm Tips
|
||||||
|
|
||||||
|
You can enable experimental memory efficient attention on pytorch 2.5 in ComfyUI on RDNA3 and potentially other AMD GPUs using this command:
|
||||||
|
|
||||||
|
```TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1 python main.py --use-pytorch-cross-attention```
|
||||||
|
|
||||||
|
You can also try setting this env variable `PYTORCH_TUNABLEOP_ENABLED=1` which might speed things up at the cost of a very slow initial run.
|
||||||
|
|
||||||
# Notes
|
# Notes
|
||||||
|
|
||||||
Only parts of the graph that have an output with all the correct inputs will be executed.
|
Only parts of the graph that have an output with all the correct inputs will be executed.
|
||||||
|
|||||||
@ -10,7 +10,6 @@ class InternalRoutes:
|
|||||||
The top level web router for internal routes: /internal/*
|
The top level web router for internal routes: /internal/*
|
||||||
The endpoints here should NOT be depended upon. It is for ComfyUI frontend use only.
|
The endpoints here should NOT be depended upon. It is for ComfyUI frontend use only.
|
||||||
Check README.md for more information.
|
Check README.md for more information.
|
||||||
|
|
||||||
'''
|
'''
|
||||||
|
|
||||||
def __init__(self, prompt_server):
|
def __init__(self, prompt_server):
|
||||||
|
|||||||
@ -36,7 +36,7 @@ class UserManager():
|
|||||||
|
|
||||||
self.settings = AppSettings(self)
|
self.settings = AppSettings(self)
|
||||||
if not os.path.exists(user_directory):
|
if not os.path.exists(user_directory):
|
||||||
os.mkdir(user_directory)
|
os.makedirs(user_directory, exist_ok=True)
|
||||||
if not args.multi_user:
|
if not args.multi_user:
|
||||||
print("****** User settings have been changed to be stored on the server instead of browser storage. ******")
|
print("****** User settings have been changed to be stored on the server instead of browser storage. ******")
|
||||||
print("****** For multi-user setups add the --multi-user CLI argument to enable multiple user profiles. ******")
|
print("****** For multi-user setups add the --multi-user CLI argument to enable multiple user profiles. ******")
|
||||||
|
|||||||
@ -2,11 +2,9 @@
|
|||||||
#and modified
|
#and modified
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch as th
|
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|
||||||
from ..ldm.modules.diffusionmodules.util import (
|
from ..ldm.modules.diffusionmodules.util import (
|
||||||
zero_module,
|
|
||||||
timestep_embedding,
|
timestep_embedding,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
120
comfy/cldm/dit_embedder.py
Normal file
120
comfy/cldm/dit_embedder.py
Normal file
@ -0,0 +1,120 @@
|
|||||||
|
import math
|
||||||
|
from typing import List, Optional, Tuple
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from torch import Tensor
|
||||||
|
|
||||||
|
from comfy.ldm.modules.diffusionmodules.mmdit import DismantledBlock, PatchEmbed, VectorEmbedder, TimestepEmbedder, get_2d_sincos_pos_embed_torch
|
||||||
|
|
||||||
|
|
||||||
|
class ControlNetEmbedder(nn.Module):
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
img_size: int,
|
||||||
|
patch_size: int,
|
||||||
|
in_chans: int,
|
||||||
|
attention_head_dim: int,
|
||||||
|
num_attention_heads: int,
|
||||||
|
adm_in_channels: int,
|
||||||
|
num_layers: int,
|
||||||
|
main_model_double: int,
|
||||||
|
double_y_emb: bool,
|
||||||
|
device: torch.device,
|
||||||
|
dtype: torch.dtype,
|
||||||
|
pos_embed_max_size: Optional[int] = None,
|
||||||
|
operations = None,
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
self.main_model_double = main_model_double
|
||||||
|
self.dtype = dtype
|
||||||
|
self.hidden_size = num_attention_heads * attention_head_dim
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.x_embedder = PatchEmbed(
|
||||||
|
img_size=img_size,
|
||||||
|
patch_size=patch_size,
|
||||||
|
in_chans=in_chans,
|
||||||
|
embed_dim=self.hidden_size,
|
||||||
|
strict_img_size=pos_embed_max_size is None,
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
operations=operations,
|
||||||
|
)
|
||||||
|
|
||||||
|
self.t_embedder = TimestepEmbedder(self.hidden_size, dtype=dtype, device=device, operations=operations)
|
||||||
|
|
||||||
|
self.double_y_emb = double_y_emb
|
||||||
|
if self.double_y_emb:
|
||||||
|
self.orig_y_embedder = VectorEmbedder(
|
||||||
|
adm_in_channels, self.hidden_size, dtype, device, operations=operations
|
||||||
|
)
|
||||||
|
self.y_embedder = VectorEmbedder(
|
||||||
|
self.hidden_size, self.hidden_size, dtype, device, operations=operations
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.y_embedder = VectorEmbedder(
|
||||||
|
adm_in_channels, self.hidden_size, dtype, device, operations=operations
|
||||||
|
)
|
||||||
|
|
||||||
|
self.transformer_blocks = nn.ModuleList(
|
||||||
|
DismantledBlock(
|
||||||
|
hidden_size=self.hidden_size, num_heads=num_attention_heads, qkv_bias=True,
|
||||||
|
dtype=dtype, device=device, operations=operations
|
||||||
|
)
|
||||||
|
for _ in range(num_layers)
|
||||||
|
)
|
||||||
|
|
||||||
|
# self.use_y_embedder = pooled_projection_dim != self.time_text_embed.text_embedder.linear_1.in_features
|
||||||
|
# TODO double check this logic when 8b
|
||||||
|
self.use_y_embedder = True
|
||||||
|
|
||||||
|
self.controlnet_blocks = nn.ModuleList([])
|
||||||
|
for _ in range(len(self.transformer_blocks)):
|
||||||
|
controlnet_block = operations.Linear(self.hidden_size, self.hidden_size, dtype=dtype, device=device)
|
||||||
|
self.controlnet_blocks.append(controlnet_block)
|
||||||
|
|
||||||
|
self.pos_embed_input = PatchEmbed(
|
||||||
|
img_size=img_size,
|
||||||
|
patch_size=patch_size,
|
||||||
|
in_chans=in_chans,
|
||||||
|
embed_dim=self.hidden_size,
|
||||||
|
strict_img_size=False,
|
||||||
|
device=device,
|
||||||
|
dtype=dtype,
|
||||||
|
operations=operations,
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x: torch.Tensor,
|
||||||
|
timesteps: torch.Tensor,
|
||||||
|
y: Optional[torch.Tensor] = None,
|
||||||
|
context: Optional[torch.Tensor] = None,
|
||||||
|
hint = None,
|
||||||
|
) -> Tuple[Tensor, List[Tensor]]:
|
||||||
|
x_shape = list(x.shape)
|
||||||
|
x = self.x_embedder(x)
|
||||||
|
if not self.double_y_emb:
|
||||||
|
h = (x_shape[-2] + 1) // self.patch_size
|
||||||
|
w = (x_shape[-1] + 1) // self.patch_size
|
||||||
|
x += get_2d_sincos_pos_embed_torch(self.hidden_size, w, h, device=x.device)
|
||||||
|
c = self.t_embedder(timesteps, dtype=x.dtype)
|
||||||
|
if y is not None and self.y_embedder is not None:
|
||||||
|
if self.double_y_emb:
|
||||||
|
y = self.orig_y_embedder(y)
|
||||||
|
y = self.y_embedder(y)
|
||||||
|
c = c + y
|
||||||
|
|
||||||
|
x = x + self.pos_embed_input(hint)
|
||||||
|
|
||||||
|
block_out = ()
|
||||||
|
|
||||||
|
repeat = math.ceil(self.main_model_double / len(self.transformer_blocks))
|
||||||
|
for i in range(len(self.transformer_blocks)):
|
||||||
|
out = self.transformer_blocks[i](x, c)
|
||||||
|
if not self.double_y_emb:
|
||||||
|
x = out
|
||||||
|
block_out += (self.controlnet_blocks[i](out),) * repeat
|
||||||
|
|
||||||
|
return {"output": block_out}
|
||||||
@ -1,5 +1,5 @@
|
|||||||
import torch
|
import torch
|
||||||
from typing import Dict, Optional
|
from typing import Optional
|
||||||
import comfy.ldm.modules.diffusionmodules.mmdit
|
import comfy.ldm.modules.diffusionmodules.mmdit
|
||||||
|
|
||||||
class ControlNet(comfy.ldm.modules.diffusionmodules.mmdit.MMDiT):
|
class ControlNet(comfy.ldm.modules.diffusionmodules.mmdit.MMDiT):
|
||||||
|
|||||||
@ -60,8 +60,10 @@ fp_group.add_argument("--force-fp32", action="store_true", help="Force fp32 (If
|
|||||||
fp_group.add_argument("--force-fp16", action="store_true", help="Force fp16.")
|
fp_group.add_argument("--force-fp16", action="store_true", help="Force fp16.")
|
||||||
|
|
||||||
fpunet_group = parser.add_mutually_exclusive_group()
|
fpunet_group = parser.add_mutually_exclusive_group()
|
||||||
fpunet_group.add_argument("--bf16-unet", action="store_true", help="Run the UNET in bf16. This should only be used for testing stuff.")
|
fpunet_group.add_argument("--fp32-unet", action="store_true", help="Run the diffusion model in fp32.")
|
||||||
fpunet_group.add_argument("--fp16-unet", action="store_true", help="Store unet weights in fp16.")
|
fpunet_group.add_argument("--fp64-unet", action="store_true", help="Run the diffusion model in fp64.")
|
||||||
|
fpunet_group.add_argument("--bf16-unet", action="store_true", help="Run the diffusion model in bf16.")
|
||||||
|
fpunet_group.add_argument("--fp16-unet", action="store_true", help="Run the diffusion model in fp16")
|
||||||
fpunet_group.add_argument("--fp8_e4m3fn-unet", action="store_true", help="Store unet weights in fp8_e4m3fn.")
|
fpunet_group.add_argument("--fp8_e4m3fn-unet", action="store_true", help="Store unet weights in fp8_e4m3fn.")
|
||||||
fpunet_group.add_argument("--fp8_e5m2-unet", action="store_true", help="Store unet weights in fp8_e5m2.")
|
fpunet_group.add_argument("--fp8_e5m2-unet", action="store_true", help="Store unet weights in fp8_e5m2.")
|
||||||
|
|
||||||
|
|||||||
@ -16,13 +16,18 @@ class Output:
|
|||||||
def __setitem__(self, key, item):
|
def __setitem__(self, key, item):
|
||||||
setattr(self, key, item)
|
setattr(self, key, item)
|
||||||
|
|
||||||
def clip_preprocess(image, size=224, mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711]):
|
def clip_preprocess(image, size=224, mean=[0.48145466, 0.4578275, 0.40821073], std=[0.26862954, 0.26130258, 0.27577711], crop=True):
|
||||||
mean = torch.tensor(mean, device=image.device, dtype=image.dtype)
|
mean = torch.tensor(mean, device=image.device, dtype=image.dtype)
|
||||||
std = torch.tensor(std, device=image.device, dtype=image.dtype)
|
std = torch.tensor(std, device=image.device, dtype=image.dtype)
|
||||||
image = image.movedim(-1, 1)
|
image = image.movedim(-1, 1)
|
||||||
if not (image.shape[2] == size and image.shape[3] == size):
|
if not (image.shape[2] == size and image.shape[3] == size):
|
||||||
scale = (size / min(image.shape[2], image.shape[3]))
|
if crop:
|
||||||
image = torch.nn.functional.interpolate(image, size=(round(scale * image.shape[2]), round(scale * image.shape[3])), mode="bicubic", antialias=True)
|
scale = (size / min(image.shape[2], image.shape[3]))
|
||||||
|
scale_size = (round(scale * image.shape[2]), round(scale * image.shape[3]))
|
||||||
|
else:
|
||||||
|
scale_size = (size, size)
|
||||||
|
|
||||||
|
image = torch.nn.functional.interpolate(image, size=scale_size, mode="bicubic", antialias=True)
|
||||||
h = (image.shape[2] - size)//2
|
h = (image.shape[2] - size)//2
|
||||||
w = (image.shape[3] - size)//2
|
w = (image.shape[3] - size)//2
|
||||||
image = image[:,:,h:h+size,w:w+size]
|
image = image[:,:,h:h+size,w:w+size]
|
||||||
@ -51,9 +56,9 @@ class ClipVisionModel():
|
|||||||
def get_sd(self):
|
def get_sd(self):
|
||||||
return self.model.state_dict()
|
return self.model.state_dict()
|
||||||
|
|
||||||
def encode_image(self, image):
|
def encode_image(self, image, crop=True):
|
||||||
comfy.model_management.load_model_gpu(self.patcher)
|
comfy.model_management.load_model_gpu(self.patcher)
|
||||||
pixel_values = clip_preprocess(image.to(self.load_device), size=self.image_size, mean=self.image_mean, std=self.image_std).float()
|
pixel_values = clip_preprocess(image.to(self.load_device), size=self.image_size, mean=self.image_mean, std=self.image_std, crop=crop).float()
|
||||||
out = self.model(pixel_values=pixel_values, intermediate_output=-2)
|
out = self.model(pixel_values=pixel_values, intermediate_output=-2)
|
||||||
|
|
||||||
outputs = Output()
|
outputs = Output()
|
||||||
|
|||||||
43
comfy/comfy_types/README.md
Normal file
43
comfy/comfy_types/README.md
Normal file
@ -0,0 +1,43 @@
|
|||||||
|
# Comfy Typing
|
||||||
|
## Type hinting for ComfyUI Node development
|
||||||
|
|
||||||
|
This module provides type hinting and concrete convenience types for node developers.
|
||||||
|
If cloned to the custom_nodes directory of ComfyUI, types can be imported using:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from comfy_types import IO, ComfyNodeABC, CheckLazyMixin
|
||||||
|
|
||||||
|
class ExampleNode(ComfyNodeABC):
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s) -> InputTypeDict:
|
||||||
|
return {"required": {}}
|
||||||
|
```
|
||||||
|
|
||||||
|
Full example is in [examples/example_nodes.py](examples/example_nodes.py).
|
||||||
|
|
||||||
|
# Types
|
||||||
|
A few primary types are documented below. More complete information is available via the docstrings on each type.
|
||||||
|
|
||||||
|
## `IO`
|
||||||
|
|
||||||
|
A string enum of built-in and a few custom data types. Includes the following special types and their requisite plumbing:
|
||||||
|
|
||||||
|
- `ANY`: `"*"`
|
||||||
|
- `NUMBER`: `"FLOAT,INT"`
|
||||||
|
- `PRIMITIVE`: `"STRING,FLOAT,INT,BOOLEAN"`
|
||||||
|
|
||||||
|
## `ComfyNodeABC`
|
||||||
|
|
||||||
|
An abstract base class for nodes, offering type-hinting / autocomplete, and somewhat-alright docstrings.
|
||||||
|
|
||||||
|
### Type hinting for `INPUT_TYPES`
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
### `INPUT_TYPES` return dict
|
||||||
|
|
||||||
|

|
||||||
|
|
||||||
|
### Options for individual inputs
|
||||||
|
|
||||||
|

|
||||||
@ -1,5 +1,6 @@
|
|||||||
import torch
|
import torch
|
||||||
from typing import Callable, Protocol, TypedDict, Optional, List
|
from typing import Callable, Protocol, TypedDict, Optional, List
|
||||||
|
from .node_typing import IO, InputTypeDict, ComfyNodeABC, CheckLazyMixin
|
||||||
|
|
||||||
|
|
||||||
class UnetApplyFunction(Protocol):
|
class UnetApplyFunction(Protocol):
|
||||||
@ -30,3 +31,15 @@ class UnetParams(TypedDict):
|
|||||||
|
|
||||||
|
|
||||||
UnetWrapperFunction = Callable[[UnetApplyFunction, UnetParams], torch.Tensor]
|
UnetWrapperFunction = Callable[[UnetApplyFunction, UnetParams], torch.Tensor]
|
||||||
|
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"UnetWrapperFunction",
|
||||||
|
UnetApplyConds.__name__,
|
||||||
|
UnetParams.__name__,
|
||||||
|
UnetApplyFunction.__name__,
|
||||||
|
IO.__name__,
|
||||||
|
InputTypeDict.__name__,
|
||||||
|
ComfyNodeABC.__name__,
|
||||||
|
CheckLazyMixin.__name__,
|
||||||
|
]
|
||||||
28
comfy/comfy_types/examples/example_nodes.py
Normal file
28
comfy/comfy_types/examples/example_nodes.py
Normal file
@ -0,0 +1,28 @@
|
|||||||
|
from comfy_types import IO, ComfyNodeABC, InputTypeDict
|
||||||
|
from inspect import cleandoc
|
||||||
|
|
||||||
|
|
||||||
|
class ExampleNode(ComfyNodeABC):
|
||||||
|
"""An example node that just adds 1 to an input integer.
|
||||||
|
|
||||||
|
* Requires an IDE configured with analysis paths etc to be worth looking at.
|
||||||
|
* Not intended for use in ComfyUI.
|
||||||
|
"""
|
||||||
|
|
||||||
|
DESCRIPTION = cleandoc(__doc__)
|
||||||
|
CATEGORY = "examples"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s) -> InputTypeDict:
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"input_int": (IO.INT, {"defaultInput": True}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = (IO.INT,)
|
||||||
|
RETURN_NAMES = ("input_plus_one",)
|
||||||
|
FUNCTION = "execute"
|
||||||
|
|
||||||
|
def execute(self, input_int: int):
|
||||||
|
return (input_int + 1,)
|
||||||
BIN
comfy/comfy_types/examples/input_options.png
Normal file
BIN
comfy/comfy_types/examples/input_options.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 19 KiB |
BIN
comfy/comfy_types/examples/input_types.png
Normal file
BIN
comfy/comfy_types/examples/input_types.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 16 KiB |
BIN
comfy/comfy_types/examples/required_hint.png
Normal file
BIN
comfy/comfy_types/examples/required_hint.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 19 KiB |
274
comfy/comfy_types/node_typing.py
Normal file
274
comfy/comfy_types/node_typing.py
Normal file
@ -0,0 +1,274 @@
|
|||||||
|
"""Comfy-specific type hinting"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
from typing import Literal, TypedDict
|
||||||
|
from abc import ABC, abstractmethod
|
||||||
|
from enum import Enum
|
||||||
|
|
||||||
|
|
||||||
|
class StrEnum(str, Enum):
|
||||||
|
"""Base class for string enums. Python's StrEnum is not available until 3.11."""
|
||||||
|
|
||||||
|
def __str__(self) -> str:
|
||||||
|
return self.value
|
||||||
|
|
||||||
|
|
||||||
|
class IO(StrEnum):
|
||||||
|
"""Node input/output data types.
|
||||||
|
|
||||||
|
Includes functionality for ``"*"`` (`ANY`) and ``"MULTI,TYPES"``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
STRING = "STRING"
|
||||||
|
IMAGE = "IMAGE"
|
||||||
|
MASK = "MASK"
|
||||||
|
LATENT = "LATENT"
|
||||||
|
BOOLEAN = "BOOLEAN"
|
||||||
|
INT = "INT"
|
||||||
|
FLOAT = "FLOAT"
|
||||||
|
CONDITIONING = "CONDITIONING"
|
||||||
|
SAMPLER = "SAMPLER"
|
||||||
|
SIGMAS = "SIGMAS"
|
||||||
|
GUIDER = "GUIDER"
|
||||||
|
NOISE = "NOISE"
|
||||||
|
CLIP = "CLIP"
|
||||||
|
CONTROL_NET = "CONTROL_NET"
|
||||||
|
VAE = "VAE"
|
||||||
|
MODEL = "MODEL"
|
||||||
|
CLIP_VISION = "CLIP_VISION"
|
||||||
|
CLIP_VISION_OUTPUT = "CLIP_VISION_OUTPUT"
|
||||||
|
STYLE_MODEL = "STYLE_MODEL"
|
||||||
|
GLIGEN = "GLIGEN"
|
||||||
|
UPSCALE_MODEL = "UPSCALE_MODEL"
|
||||||
|
AUDIO = "AUDIO"
|
||||||
|
WEBCAM = "WEBCAM"
|
||||||
|
POINT = "POINT"
|
||||||
|
FACE_ANALYSIS = "FACE_ANALYSIS"
|
||||||
|
BBOX = "BBOX"
|
||||||
|
SEGS = "SEGS"
|
||||||
|
|
||||||
|
ANY = "*"
|
||||||
|
"""Always matches any type, but at a price.
|
||||||
|
|
||||||
|
Causes some functionality issues (e.g. reroutes, link types), and should be avoided whenever possible.
|
||||||
|
"""
|
||||||
|
NUMBER = "FLOAT,INT"
|
||||||
|
"""A float or an int - could be either"""
|
||||||
|
PRIMITIVE = "STRING,FLOAT,INT,BOOLEAN"
|
||||||
|
"""Could be any of: string, float, int, or bool"""
|
||||||
|
|
||||||
|
def __ne__(self, value: object) -> bool:
|
||||||
|
if self == "*" or value == "*":
|
||||||
|
return False
|
||||||
|
if not isinstance(value, str):
|
||||||
|
return True
|
||||||
|
a = frozenset(self.split(","))
|
||||||
|
b = frozenset(value.split(","))
|
||||||
|
return not (b.issubset(a) or a.issubset(b))
|
||||||
|
|
||||||
|
|
||||||
|
class InputTypeOptions(TypedDict):
|
||||||
|
"""Provides type hinting for the return type of the INPUT_TYPES node function.
|
||||||
|
|
||||||
|
Due to IDE limitations with unions, for now all options are available for all types (e.g. `label_on` is hinted even when the type is not `IO.BOOLEAN`).
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_datatypes
|
||||||
|
"""
|
||||||
|
|
||||||
|
default: bool | str | float | int | list | tuple
|
||||||
|
"""The default value of the widget"""
|
||||||
|
defaultInput: bool
|
||||||
|
"""Defaults to an input slot rather than a widget"""
|
||||||
|
forceInput: bool
|
||||||
|
"""`defaultInput` and also don't allow converting to a widget"""
|
||||||
|
lazy: bool
|
||||||
|
"""Declares that this input uses lazy evaluation"""
|
||||||
|
rawLink: bool
|
||||||
|
"""When a link exists, rather than receiving the evaluated value, you will receive the link (i.e. `["nodeId", <outputIndex>]`). Designed for node expansion."""
|
||||||
|
tooltip: str
|
||||||
|
"""Tooltip for the input (or widget), shown on pointer hover"""
|
||||||
|
# class InputTypeNumber(InputTypeOptions):
|
||||||
|
# default: float | int
|
||||||
|
min: float
|
||||||
|
"""The minimum value of a number (``FLOAT`` | ``INT``)"""
|
||||||
|
max: float
|
||||||
|
"""The maximum value of a number (``FLOAT`` | ``INT``)"""
|
||||||
|
step: float
|
||||||
|
"""The amount to increment or decrement a widget by when stepping up/down (``FLOAT`` | ``INT``)"""
|
||||||
|
round: float
|
||||||
|
"""Floats are rounded by this value (``FLOAT``)"""
|
||||||
|
# class InputTypeBoolean(InputTypeOptions):
|
||||||
|
# default: bool
|
||||||
|
label_on: str
|
||||||
|
"""The label to use in the UI when the bool is True (``BOOLEAN``)"""
|
||||||
|
label_on: str
|
||||||
|
"""The label to use in the UI when the bool is False (``BOOLEAN``)"""
|
||||||
|
# class InputTypeString(InputTypeOptions):
|
||||||
|
# default: str
|
||||||
|
multiline: bool
|
||||||
|
"""Use a multiline text box (``STRING``)"""
|
||||||
|
placeholder: str
|
||||||
|
"""Placeholder text to display in the UI when empty (``STRING``)"""
|
||||||
|
# Deprecated:
|
||||||
|
# defaultVal: str
|
||||||
|
dynamicPrompts: bool
|
||||||
|
"""Causes the front-end to evaluate dynamic prompts (``STRING``)"""
|
||||||
|
|
||||||
|
|
||||||
|
class HiddenInputTypeDict(TypedDict):
|
||||||
|
"""Provides type hinting for the hidden entry of node INPUT_TYPES."""
|
||||||
|
|
||||||
|
node_id: Literal["UNIQUE_ID"]
|
||||||
|
"""UNIQUE_ID is the unique identifier of the node, and matches the id property of the node on the client side. It is commonly used in client-server communications (see messages)."""
|
||||||
|
unique_id: Literal["UNIQUE_ID"]
|
||||||
|
"""UNIQUE_ID is the unique identifier of the node, and matches the id property of the node on the client side. It is commonly used in client-server communications (see messages)."""
|
||||||
|
prompt: Literal["PROMPT"]
|
||||||
|
"""PROMPT is the complete prompt sent by the client to the server. See the prompt object for a full description."""
|
||||||
|
extra_pnginfo: Literal["EXTRA_PNGINFO"]
|
||||||
|
"""EXTRA_PNGINFO is a dictionary that will be copied into the metadata of any .png files saved. Custom nodes can store additional information in this dictionary for saving (or as a way to communicate with a downstream node)."""
|
||||||
|
dynprompt: Literal["DYNPROMPT"]
|
||||||
|
"""DYNPROMPT is an instance of comfy_execution.graph.DynamicPrompt. It differs from PROMPT in that it may mutate during the course of execution in response to Node Expansion."""
|
||||||
|
|
||||||
|
|
||||||
|
class InputTypeDict(TypedDict):
|
||||||
|
"""Provides type hinting for node INPUT_TYPES.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_more_on_inputs
|
||||||
|
"""
|
||||||
|
|
||||||
|
required: dict[str, tuple[IO, InputTypeOptions]]
|
||||||
|
"""Describes all inputs that must be connected for the node to execute."""
|
||||||
|
optional: dict[str, tuple[IO, InputTypeOptions]]
|
||||||
|
"""Describes inputs which do not need to be connected."""
|
||||||
|
hidden: HiddenInputTypeDict
|
||||||
|
"""Offers advanced functionality and server-client communication.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_more_on_inputs#hidden-inputs
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class ComfyNodeABC(ABC):
|
||||||
|
"""Abstract base class for Comfy nodes. Includes the names and expected types of attributes.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview
|
||||||
|
"""
|
||||||
|
|
||||||
|
DESCRIPTION: str
|
||||||
|
"""Node description, shown as a tooltip when hovering over the node.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
# Explicitly define the description
|
||||||
|
DESCRIPTION = "Example description here."
|
||||||
|
|
||||||
|
# Use the docstring of the node class.
|
||||||
|
DESCRIPTION = cleandoc(__doc__)
|
||||||
|
"""
|
||||||
|
CATEGORY: str
|
||||||
|
"""The category of the node, as per the "Add Node" menu.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview#category
|
||||||
|
"""
|
||||||
|
EXPERIMENTAL: bool
|
||||||
|
"""Flags a node as experimental, informing users that it may change or not work as expected."""
|
||||||
|
DEPRECATED: bool
|
||||||
|
"""Flags a node as deprecated, indicating to users that they should find alternatives to this node."""
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
@abstractmethod
|
||||||
|
def INPUT_TYPES(s) -> InputTypeDict:
|
||||||
|
"""Defines node inputs.
|
||||||
|
|
||||||
|
* Must include the ``required`` key, which describes all inputs that must be connected for the node to execute.
|
||||||
|
* The ``optional`` key can be added to describe inputs which do not need to be connected.
|
||||||
|
* The ``hidden`` key offers some advanced functionality. More info at: https://docs.comfy.org/essentials/custom_node_more_on_inputs#hidden-inputs
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview#input-types
|
||||||
|
"""
|
||||||
|
return {"required": {}}
|
||||||
|
|
||||||
|
OUTPUT_NODE: bool
|
||||||
|
"""Flags this node as an output node, causing any inputs it requires to be executed.
|
||||||
|
|
||||||
|
If a node is not connected to any output nodes, that node will not be executed. Usage::
|
||||||
|
|
||||||
|
OUTPUT_NODE = True
|
||||||
|
|
||||||
|
From the docs:
|
||||||
|
|
||||||
|
By default, a node is not considered an output. Set ``OUTPUT_NODE = True`` to specify that it is.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview#output-node
|
||||||
|
"""
|
||||||
|
INPUT_IS_LIST: bool
|
||||||
|
"""A flag indicating if this node implements the additional code necessary to deal with OUTPUT_IS_LIST nodes.
|
||||||
|
|
||||||
|
All inputs of ``type`` will become ``list[type]``, regardless of how many items are passed in. This also affects ``check_lazy_status``.
|
||||||
|
|
||||||
|
From the docs:
|
||||||
|
|
||||||
|
A node can also override the default input behaviour and receive the whole list in a single call. This is done by setting a class attribute `INPUT_IS_LIST` to ``True``.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_lists#list-processing
|
||||||
|
"""
|
||||||
|
OUTPUT_IS_LIST: tuple[bool]
|
||||||
|
"""A tuple indicating which node outputs are lists, but will be connected to nodes that expect individual items.
|
||||||
|
|
||||||
|
Connected nodes that do not implement `INPUT_IS_LIST` will be executed once for every item in the list.
|
||||||
|
|
||||||
|
A ``tuple[bool]``, where the items match those in `RETURN_TYPES`::
|
||||||
|
|
||||||
|
RETURN_TYPES = (IO.INT, IO.INT, IO.STRING)
|
||||||
|
OUTPUT_IS_LIST = (True, True, False) # The string output will be handled normally
|
||||||
|
|
||||||
|
From the docs:
|
||||||
|
|
||||||
|
In order to tell Comfy that the list being returned should not be wrapped, but treated as a series of data for sequential processing,
|
||||||
|
the node should provide a class attribute `OUTPUT_IS_LIST`, which is a ``tuple[bool]``, of the same length as `RETURN_TYPES`,
|
||||||
|
specifying which outputs which should be so treated.
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_lists#list-processing
|
||||||
|
"""
|
||||||
|
|
||||||
|
RETURN_TYPES: tuple[IO]
|
||||||
|
"""A tuple representing the outputs of this node.
|
||||||
|
|
||||||
|
Usage::
|
||||||
|
|
||||||
|
RETURN_TYPES = (IO.INT, "INT", "CUSTOM_TYPE")
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview#return-types
|
||||||
|
"""
|
||||||
|
RETURN_NAMES: tuple[str]
|
||||||
|
"""The output slot names for each item in `RETURN_TYPES`, e.g. ``RETURN_NAMES = ("count", "filter_string")``
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview#return-names
|
||||||
|
"""
|
||||||
|
OUTPUT_TOOLTIPS: tuple[str]
|
||||||
|
"""A tuple of strings to use as tooltips for node outputs, one for each item in `RETURN_TYPES`."""
|
||||||
|
FUNCTION: str
|
||||||
|
"""The name of the function to execute as a literal string, e.g. `FUNCTION = "execute"`
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_server_overview#function
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
class CheckLazyMixin:
|
||||||
|
"""Provides a basic check_lazy_status implementation and type hinting for nodes that use lazy inputs."""
|
||||||
|
|
||||||
|
def check_lazy_status(self, **kwargs) -> list[str]:
|
||||||
|
"""Returns a list of input names that should be evaluated.
|
||||||
|
|
||||||
|
This basic mixin impl. requires all inputs.
|
||||||
|
|
||||||
|
:kwargs: All node inputs will be included here. If the input is ``None``, it should be assumed that it has not yet been evaluated. \
|
||||||
|
When using ``INPUT_IS_LIST = True``, unevaluated will instead be ``(None,)``.
|
||||||
|
|
||||||
|
Params should match the nodes execution ``FUNCTION`` (self, and all inputs by name).
|
||||||
|
Will be executed repeatedly until it returns an empty list, or all requested items were already evaluated (and sent as params).
|
||||||
|
|
||||||
|
Comfy Docs: https://docs.comfy.org/essentials/custom_node_lazy_evaluation#defining-check-lazy-status
|
||||||
|
"""
|
||||||
|
|
||||||
|
need = [name for name in kwargs if kwargs[name] is None]
|
||||||
|
return need
|
||||||
@ -35,6 +35,10 @@ import comfy.ldm.cascade.controlnet
|
|||||||
import comfy.cldm.mmdit
|
import comfy.cldm.mmdit
|
||||||
import comfy.ldm.hydit.controlnet
|
import comfy.ldm.hydit.controlnet
|
||||||
import comfy.ldm.flux.controlnet
|
import comfy.ldm.flux.controlnet
|
||||||
|
import comfy.cldm.dit_embedder
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from comfy.hooks import HookGroup
|
||||||
|
|
||||||
|
|
||||||
def broadcast_image_to(tensor, target_batch_size, batched_number):
|
def broadcast_image_to(tensor, target_batch_size, batched_number):
|
||||||
@ -78,6 +82,8 @@ class ControlBase:
|
|||||||
self.concat_mask = False
|
self.concat_mask = False
|
||||||
self.extra_concat_orig = []
|
self.extra_concat_orig = []
|
||||||
self.extra_concat = None
|
self.extra_concat = None
|
||||||
|
self.extra_hooks: HookGroup = None
|
||||||
|
self.preprocess_image = lambda a: a
|
||||||
|
|
||||||
def set_cond_hint(self, cond_hint, strength=1.0, timestep_percent_range=(0.0, 1.0), vae=None, extra_concat=[]):
|
def set_cond_hint(self, cond_hint, strength=1.0, timestep_percent_range=(0.0, 1.0), vae=None, extra_concat=[]):
|
||||||
self.cond_hint_original = cond_hint
|
self.cond_hint_original = cond_hint
|
||||||
@ -114,6 +120,14 @@ class ControlBase:
|
|||||||
if self.previous_controlnet is not None:
|
if self.previous_controlnet is not None:
|
||||||
out += self.previous_controlnet.get_models()
|
out += self.previous_controlnet.get_models()
|
||||||
return out
|
return out
|
||||||
|
|
||||||
|
def get_extra_hooks(self):
|
||||||
|
out = []
|
||||||
|
if self.extra_hooks is not None:
|
||||||
|
out.append(self.extra_hooks)
|
||||||
|
if self.previous_controlnet is not None:
|
||||||
|
out += self.previous_controlnet.get_extra_hooks()
|
||||||
|
return out
|
||||||
|
|
||||||
def copy_to(self, c):
|
def copy_to(self, c):
|
||||||
c.cond_hint_original = self.cond_hint_original
|
c.cond_hint_original = self.cond_hint_original
|
||||||
@ -129,6 +143,8 @@ class ControlBase:
|
|||||||
c.strength_type = self.strength_type
|
c.strength_type = self.strength_type
|
||||||
c.concat_mask = self.concat_mask
|
c.concat_mask = self.concat_mask
|
||||||
c.extra_concat_orig = self.extra_concat_orig.copy()
|
c.extra_concat_orig = self.extra_concat_orig.copy()
|
||||||
|
c.extra_hooks = self.extra_hooks.clone() if self.extra_hooks else None
|
||||||
|
c.preprocess_image = self.preprocess_image
|
||||||
|
|
||||||
def inference_memory_requirements(self, dtype):
|
def inference_memory_requirements(self, dtype):
|
||||||
if self.previous_controlnet is not None:
|
if self.previous_controlnet is not None:
|
||||||
@ -181,7 +197,7 @@ class ControlBase:
|
|||||||
|
|
||||||
|
|
||||||
class ControlNet(ControlBase):
|
class ControlNet(ControlBase):
|
||||||
def __init__(self, control_model=None, global_average_pooling=False, compression_ratio=8, latent_format=None, load_device=None, manual_cast_dtype=None, extra_conds=["y"], strength_type=StrengthType.CONSTANT, concat_mask=False):
|
def __init__(self, control_model=None, global_average_pooling=False, compression_ratio=8, latent_format=None, load_device=None, manual_cast_dtype=None, extra_conds=["y"], strength_type=StrengthType.CONSTANT, concat_mask=False, preprocess_image=lambda a: a):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.control_model = control_model
|
self.control_model = control_model
|
||||||
self.load_device = load_device
|
self.load_device = load_device
|
||||||
@ -196,11 +212,12 @@ class ControlNet(ControlBase):
|
|||||||
self.extra_conds += extra_conds
|
self.extra_conds += extra_conds
|
||||||
self.strength_type = strength_type
|
self.strength_type = strength_type
|
||||||
self.concat_mask = concat_mask
|
self.concat_mask = concat_mask
|
||||||
|
self.preprocess_image = preprocess_image
|
||||||
|
|
||||||
def get_control(self, x_noisy, t, cond, batched_number):
|
def get_control(self, x_noisy, t, cond, batched_number, transformer_options):
|
||||||
control_prev = None
|
control_prev = None
|
||||||
if self.previous_controlnet is not None:
|
if self.previous_controlnet is not None:
|
||||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number, transformer_options)
|
||||||
|
|
||||||
if self.timestep_range is not None:
|
if self.timestep_range is not None:
|
||||||
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
||||||
@ -224,6 +241,7 @@ class ControlNet(ControlBase):
|
|||||||
if self.latent_format is not None:
|
if self.latent_format is not None:
|
||||||
raise ValueError("This Controlnet needs a VAE but none was provided, please use a ControlNetApply node with a VAE input and connect it.")
|
raise ValueError("This Controlnet needs a VAE but none was provided, please use a ControlNetApply node with a VAE input and connect it.")
|
||||||
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, self.upscale_algorithm, "center")
|
self.cond_hint = comfy.utils.common_upscale(self.cond_hint_original, x_noisy.shape[3] * compression_ratio, x_noisy.shape[2] * compression_ratio, self.upscale_algorithm, "center")
|
||||||
|
self.cond_hint = self.preprocess_image(self.cond_hint)
|
||||||
if self.vae is not None:
|
if self.vae is not None:
|
||||||
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
|
loaded_models = comfy.model_management.loaded_models(only_currently_used=True)
|
||||||
self.cond_hint = self.vae.encode(self.cond_hint.movedim(1, -1))
|
self.cond_hint = self.vae.encode(self.cond_hint.movedim(1, -1))
|
||||||
@ -427,6 +445,7 @@ def controlnet_load_state_dict(control_model, sd):
|
|||||||
logging.debug("unexpected controlnet keys: {}".format(unexpected))
|
logging.debug("unexpected controlnet keys: {}".format(unexpected))
|
||||||
return control_model
|
return control_model
|
||||||
|
|
||||||
|
|
||||||
def load_controlnet_mmdit(sd, model_options={}):
|
def load_controlnet_mmdit(sd, model_options={}):
|
||||||
new_sd = comfy.model_detection.convert_diffusers_mmdit(sd, "")
|
new_sd = comfy.model_detection.convert_diffusers_mmdit(sd, "")
|
||||||
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(new_sd, model_options=model_options)
|
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(new_sd, model_options=model_options)
|
||||||
@ -448,6 +467,82 @@ def load_controlnet_mmdit(sd, model_options={}):
|
|||||||
return control
|
return control
|
||||||
|
|
||||||
|
|
||||||
|
class ControlNetSD35(ControlNet):
|
||||||
|
def pre_run(self, model, percent_to_timestep_function):
|
||||||
|
if self.control_model.double_y_emb:
|
||||||
|
missing, unexpected = self.control_model.orig_y_embedder.load_state_dict(model.diffusion_model.y_embedder.state_dict(), strict=False)
|
||||||
|
else:
|
||||||
|
missing, unexpected = self.control_model.x_embedder.load_state_dict(model.diffusion_model.x_embedder.state_dict(), strict=False)
|
||||||
|
super().pre_run(model, percent_to_timestep_function)
|
||||||
|
|
||||||
|
def copy(self):
|
||||||
|
c = ControlNetSD35(None, global_average_pooling=self.global_average_pooling, load_device=self.load_device, manual_cast_dtype=self.manual_cast_dtype)
|
||||||
|
c.control_model = self.control_model
|
||||||
|
c.control_model_wrapped = self.control_model_wrapped
|
||||||
|
self.copy_to(c)
|
||||||
|
return c
|
||||||
|
|
||||||
|
def load_controlnet_sd35(sd, model_options={}):
|
||||||
|
control_type = -1
|
||||||
|
if "control_type" in sd:
|
||||||
|
control_type = round(sd.pop("control_type").item())
|
||||||
|
|
||||||
|
# blur_cnet = control_type == 0
|
||||||
|
canny_cnet = control_type == 1
|
||||||
|
depth_cnet = control_type == 2
|
||||||
|
|
||||||
|
new_sd = {}
|
||||||
|
for k in comfy.utils.MMDIT_MAP_BASIC:
|
||||||
|
if k[1] in sd:
|
||||||
|
new_sd[k[0]] = sd.pop(k[1])
|
||||||
|
for k in sd:
|
||||||
|
new_sd[k] = sd[k]
|
||||||
|
sd = new_sd
|
||||||
|
|
||||||
|
y_emb_shape = sd["y_embedder.mlp.0.weight"].shape
|
||||||
|
depth = y_emb_shape[0] // 64
|
||||||
|
hidden_size = 64 * depth
|
||||||
|
num_heads = depth
|
||||||
|
head_dim = hidden_size // num_heads
|
||||||
|
num_blocks = comfy.model_detection.count_blocks(new_sd, 'transformer_blocks.{}.')
|
||||||
|
|
||||||
|
load_device = comfy.model_management.get_torch_device()
|
||||||
|
offload_device = comfy.model_management.unet_offload_device()
|
||||||
|
unet_dtype = comfy.model_management.unet_dtype(model_params=-1)
|
||||||
|
|
||||||
|
manual_cast_dtype = comfy.model_management.unet_manual_cast(unet_dtype, load_device)
|
||||||
|
|
||||||
|
operations = model_options.get("custom_operations", None)
|
||||||
|
if operations is None:
|
||||||
|
operations = comfy.ops.pick_operations(unet_dtype, manual_cast_dtype, disable_fast_fp8=True)
|
||||||
|
|
||||||
|
control_model = comfy.cldm.dit_embedder.ControlNetEmbedder(img_size=None,
|
||||||
|
patch_size=2,
|
||||||
|
in_chans=16,
|
||||||
|
num_layers=num_blocks,
|
||||||
|
main_model_double=depth,
|
||||||
|
double_y_emb=y_emb_shape[0] == y_emb_shape[1],
|
||||||
|
attention_head_dim=head_dim,
|
||||||
|
num_attention_heads=num_heads,
|
||||||
|
adm_in_channels=2048,
|
||||||
|
device=offload_device,
|
||||||
|
dtype=unet_dtype,
|
||||||
|
operations=operations)
|
||||||
|
|
||||||
|
control_model = controlnet_load_state_dict(control_model, sd)
|
||||||
|
|
||||||
|
latent_format = comfy.latent_formats.SD3()
|
||||||
|
preprocess_image = lambda a: a
|
||||||
|
if canny_cnet:
|
||||||
|
preprocess_image = lambda a: (a * 255 * 0.5 + 0.5)
|
||||||
|
elif depth_cnet:
|
||||||
|
preprocess_image = lambda a: 1.0 - a
|
||||||
|
|
||||||
|
control = ControlNetSD35(control_model, compression_ratio=1, latent_format=latent_format, load_device=load_device, manual_cast_dtype=manual_cast_dtype, preprocess_image=preprocess_image)
|
||||||
|
return control
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def load_controlnet_hunyuandit(controlnet_data, model_options={}):
|
def load_controlnet_hunyuandit(controlnet_data, model_options={}):
|
||||||
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(controlnet_data, model_options=model_options)
|
model_config, operations, load_device, unet_dtype, manual_cast_dtype, offload_device = controlnet_config(controlnet_data, model_options=model_options)
|
||||||
|
|
||||||
@ -560,7 +655,10 @@ def load_controlnet_state_dict(state_dict, model=None, model_options={}):
|
|||||||
if "double_blocks.0.img_attn.norm.key_norm.scale" in controlnet_data:
|
if "double_blocks.0.img_attn.norm.key_norm.scale" in controlnet_data:
|
||||||
return load_controlnet_flux_xlabs_mistoline(controlnet_data, model_options=model_options)
|
return load_controlnet_flux_xlabs_mistoline(controlnet_data, model_options=model_options)
|
||||||
elif "pos_embed_input.proj.weight" in controlnet_data:
|
elif "pos_embed_input.proj.weight" in controlnet_data:
|
||||||
return load_controlnet_mmdit(controlnet_data, model_options=model_options) #SD3 diffusers controlnet
|
if "transformer_blocks.0.adaLN_modulation.1.bias" in controlnet_data:
|
||||||
|
return load_controlnet_sd35(controlnet_data, model_options=model_options) #Stability sd3.5 format
|
||||||
|
else:
|
||||||
|
return load_controlnet_mmdit(controlnet_data, model_options=model_options) #SD3 diffusers controlnet
|
||||||
elif "controlnet_x_embedder.weight" in controlnet_data:
|
elif "controlnet_x_embedder.weight" in controlnet_data:
|
||||||
return load_controlnet_flux_instantx(controlnet_data, model_options=model_options)
|
return load_controlnet_flux_instantx(controlnet_data, model_options=model_options)
|
||||||
elif "controlnet_blocks.0.linear.weight" in controlnet_data: #mistoline flux
|
elif "controlnet_blocks.0.linear.weight" in controlnet_data: #mistoline flux
|
||||||
@ -674,10 +772,10 @@ class T2IAdapter(ControlBase):
|
|||||||
height = math.ceil(height / unshuffle_amount) * unshuffle_amount
|
height = math.ceil(height / unshuffle_amount) * unshuffle_amount
|
||||||
return width, height
|
return width, height
|
||||||
|
|
||||||
def get_control(self, x_noisy, t, cond, batched_number):
|
def get_control(self, x_noisy, t, cond, batched_number, transformer_options):
|
||||||
control_prev = None
|
control_prev = None
|
||||||
if self.previous_controlnet is not None:
|
if self.previous_controlnet is not None:
|
||||||
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number)
|
control_prev = self.previous_controlnet.get_control(x_noisy, t, cond, batched_number, transformer_options)
|
||||||
|
|
||||||
if self.timestep_range is not None:
|
if self.timestep_range is not None:
|
||||||
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
if t[0] > self.timestep_range[0] or t[0] < self.timestep_range[1]:
|
||||||
|
|||||||
@ -1,10 +1,9 @@
|
|||||||
#code taken from: https://github.com/wl-zhao/UniPC and modified
|
#code taken from: https://github.com/wl-zhao/UniPC and modified
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
|
||||||
import math
|
import math
|
||||||
|
|
||||||
from tqdm.auto import trange, tqdm
|
from tqdm.auto import trange
|
||||||
|
|
||||||
|
|
||||||
class NoiseScheduleVP:
|
class NoiseScheduleVP:
|
||||||
|
|||||||
690
comfy/hooks.py
Normal file
690
comfy/hooks.py
Normal file
@ -0,0 +1,690 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from typing import TYPE_CHECKING, Callable
|
||||||
|
import enum
|
||||||
|
import math
|
||||||
|
import torch
|
||||||
|
import numpy as np
|
||||||
|
import itertools
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from comfy.model_patcher import ModelPatcher, PatcherInjection
|
||||||
|
from comfy.model_base import BaseModel
|
||||||
|
from comfy.sd import CLIP
|
||||||
|
import comfy.lora
|
||||||
|
import comfy.model_management
|
||||||
|
import comfy.patcher_extension
|
||||||
|
from node_helpers import conditioning_set_values
|
||||||
|
|
||||||
|
class EnumHookMode(enum.Enum):
|
||||||
|
MinVram = "minvram"
|
||||||
|
MaxSpeed = "maxspeed"
|
||||||
|
|
||||||
|
class EnumHookType(enum.Enum):
|
||||||
|
Weight = "weight"
|
||||||
|
Patch = "patch"
|
||||||
|
ObjectPatch = "object_patch"
|
||||||
|
AddModels = "add_models"
|
||||||
|
Callbacks = "callbacks"
|
||||||
|
Wrappers = "wrappers"
|
||||||
|
SetInjections = "add_injections"
|
||||||
|
|
||||||
|
class EnumWeightTarget(enum.Enum):
|
||||||
|
Model = "model"
|
||||||
|
Clip = "clip"
|
||||||
|
|
||||||
|
class _HookRef:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# NOTE: this is an example of how the should_register function should look
|
||||||
|
def default_should_register(hook: 'Hook', model: 'ModelPatcher', model_options: dict, target: EnumWeightTarget, registered: list[Hook]):
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class Hook:
|
||||||
|
def __init__(self, hook_type: EnumHookType=None, hook_ref: _HookRef=None, hook_id: str=None,
|
||||||
|
hook_keyframe: 'HookKeyframeGroup'=None):
|
||||||
|
self.hook_type = hook_type
|
||||||
|
self.hook_ref = hook_ref if hook_ref else _HookRef()
|
||||||
|
self.hook_id = hook_id
|
||||||
|
self.hook_keyframe = hook_keyframe if hook_keyframe else HookKeyframeGroup()
|
||||||
|
self.custom_should_register = default_should_register
|
||||||
|
self.auto_apply_to_nonpositive = False
|
||||||
|
|
||||||
|
@property
|
||||||
|
def strength(self):
|
||||||
|
return self.hook_keyframe.strength
|
||||||
|
|
||||||
|
def initialize_timesteps(self, model: 'BaseModel'):
|
||||||
|
self.reset()
|
||||||
|
self.hook_keyframe.initialize_timesteps(model)
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self.hook_keyframe.reset()
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: Hook = subtype()
|
||||||
|
c.hook_type = self.hook_type
|
||||||
|
c.hook_ref = self.hook_ref
|
||||||
|
c.hook_id = self.hook_id
|
||||||
|
c.hook_keyframe = self.hook_keyframe
|
||||||
|
c.custom_should_register = self.custom_should_register
|
||||||
|
# TODO: make this do something
|
||||||
|
c.auto_apply_to_nonpositive = self.auto_apply_to_nonpositive
|
||||||
|
return c
|
||||||
|
|
||||||
|
def should_register(self, model: 'ModelPatcher', model_options: dict, target: EnumWeightTarget, registered: list[Hook]):
|
||||||
|
return self.custom_should_register(self, model, model_options, target, registered)
|
||||||
|
|
||||||
|
def add_hook_patches(self, model: 'ModelPatcher', model_options: dict, target: EnumWeightTarget, registered: list[Hook]):
|
||||||
|
raise NotImplementedError("add_hook_patches should be defined for Hook subclasses")
|
||||||
|
|
||||||
|
def on_apply(self, model: 'ModelPatcher', transformer_options: dict[str]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def on_unapply(self, model: 'ModelPatcher', transformer_options: dict[str]):
|
||||||
|
pass
|
||||||
|
|
||||||
|
def __eq__(self, other: 'Hook'):
|
||||||
|
return self.__class__ == other.__class__ and self.hook_ref == other.hook_ref
|
||||||
|
|
||||||
|
def __hash__(self):
|
||||||
|
return hash(self.hook_ref)
|
||||||
|
|
||||||
|
class WeightHook(Hook):
|
||||||
|
def __init__(self, strength_model=1.0, strength_clip=1.0):
|
||||||
|
super().__init__(hook_type=EnumHookType.Weight)
|
||||||
|
self.weights: dict = None
|
||||||
|
self.weights_clip: dict = None
|
||||||
|
self.need_weight_init = True
|
||||||
|
self._strength_model = strength_model
|
||||||
|
self._strength_clip = strength_clip
|
||||||
|
|
||||||
|
@property
|
||||||
|
def strength_model(self):
|
||||||
|
return self._strength_model * self.strength
|
||||||
|
|
||||||
|
@property
|
||||||
|
def strength_clip(self):
|
||||||
|
return self._strength_clip * self.strength
|
||||||
|
|
||||||
|
def add_hook_patches(self, model: 'ModelPatcher', model_options: dict, target: EnumWeightTarget, registered: list[Hook]):
|
||||||
|
if not self.should_register(model, model_options, target, registered):
|
||||||
|
return False
|
||||||
|
weights = None
|
||||||
|
if target == EnumWeightTarget.Model:
|
||||||
|
strength = self._strength_model
|
||||||
|
else:
|
||||||
|
strength = self._strength_clip
|
||||||
|
|
||||||
|
if self.need_weight_init:
|
||||||
|
key_map = {}
|
||||||
|
if target == EnumWeightTarget.Model:
|
||||||
|
key_map = comfy.lora.model_lora_keys_unet(model.model, key_map)
|
||||||
|
else:
|
||||||
|
key_map = comfy.lora.model_lora_keys_clip(model.model, key_map)
|
||||||
|
weights = comfy.lora.load_lora(self.weights, key_map, log_missing=False)
|
||||||
|
else:
|
||||||
|
if target == EnumWeightTarget.Model:
|
||||||
|
weights = self.weights
|
||||||
|
else:
|
||||||
|
weights = self.weights_clip
|
||||||
|
k = model.add_hook_patches(hook=self, patches=weights, strength_patch=strength)
|
||||||
|
registered.append(self)
|
||||||
|
return True
|
||||||
|
# TODO: add logs about any keys that were not applied
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: WeightHook = super().clone(subtype)
|
||||||
|
c.weights = self.weights
|
||||||
|
c.weights_clip = self.weights_clip
|
||||||
|
c.need_weight_init = self.need_weight_init
|
||||||
|
c._strength_model = self._strength_model
|
||||||
|
c._strength_clip = self._strength_clip
|
||||||
|
return c
|
||||||
|
|
||||||
|
class PatchHook(Hook):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__(hook_type=EnumHookType.Patch)
|
||||||
|
self.patches: dict = None
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: PatchHook = super().clone(subtype)
|
||||||
|
c.patches = self.patches
|
||||||
|
return c
|
||||||
|
# TODO: add functionality
|
||||||
|
|
||||||
|
class ObjectPatchHook(Hook):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__(hook_type=EnumHookType.ObjectPatch)
|
||||||
|
self.object_patches: dict = None
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: ObjectPatchHook = super().clone(subtype)
|
||||||
|
c.object_patches = self.object_patches
|
||||||
|
return c
|
||||||
|
# TODO: add functionality
|
||||||
|
|
||||||
|
class AddModelsHook(Hook):
|
||||||
|
def __init__(self, key: str=None, models: list['ModelPatcher']=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.AddModels)
|
||||||
|
self.key = key
|
||||||
|
self.models = models
|
||||||
|
self.append_when_same = True
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: AddModelsHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.models = self.models.copy() if self.models else self.models
|
||||||
|
c.append_when_same = self.append_when_same
|
||||||
|
return c
|
||||||
|
# TODO: add functionality
|
||||||
|
|
||||||
|
class CallbackHook(Hook):
|
||||||
|
def __init__(self, key: str=None, callback: Callable=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.Callbacks)
|
||||||
|
self.key = key
|
||||||
|
self.callback = callback
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: CallbackHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.callback = self.callback
|
||||||
|
return c
|
||||||
|
# TODO: add functionality
|
||||||
|
|
||||||
|
class WrapperHook(Hook):
|
||||||
|
def __init__(self, wrappers_dict: dict[str, dict[str, dict[str, list[Callable]]]]=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.Wrappers)
|
||||||
|
self.wrappers_dict = wrappers_dict
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: WrapperHook = super().clone(subtype)
|
||||||
|
c.wrappers_dict = self.wrappers_dict
|
||||||
|
return c
|
||||||
|
|
||||||
|
def add_hook_patches(self, model: 'ModelPatcher', model_options: dict, target: EnumWeightTarget, registered: list[Hook]):
|
||||||
|
if not self.should_register(model, model_options, target, registered):
|
||||||
|
return False
|
||||||
|
add_model_options = {"transformer_options": self.wrappers_dict}
|
||||||
|
comfy.patcher_extension.merge_nested_dicts(model_options, add_model_options, copy_dict1=False)
|
||||||
|
registered.append(self)
|
||||||
|
return True
|
||||||
|
|
||||||
|
class SetInjectionsHook(Hook):
|
||||||
|
def __init__(self, key: str=None, injections: list['PatcherInjection']=None):
|
||||||
|
super().__init__(hook_type=EnumHookType.SetInjections)
|
||||||
|
self.key = key
|
||||||
|
self.injections = injections
|
||||||
|
|
||||||
|
def clone(self, subtype: Callable=None):
|
||||||
|
if subtype is None:
|
||||||
|
subtype = type(self)
|
||||||
|
c: SetInjectionsHook = super().clone(subtype)
|
||||||
|
c.key = self.key
|
||||||
|
c.injections = self.injections.copy() if self.injections else self.injections
|
||||||
|
return c
|
||||||
|
|
||||||
|
def add_hook_injections(self, model: 'ModelPatcher'):
|
||||||
|
# TODO: add functionality
|
||||||
|
pass
|
||||||
|
|
||||||
|
class HookGroup:
|
||||||
|
def __init__(self):
|
||||||
|
self.hooks: list[Hook] = []
|
||||||
|
|
||||||
|
def add(self, hook: Hook):
|
||||||
|
if hook not in self.hooks:
|
||||||
|
self.hooks.append(hook)
|
||||||
|
|
||||||
|
def contains(self, hook: Hook):
|
||||||
|
return hook in self.hooks
|
||||||
|
|
||||||
|
def clone(self):
|
||||||
|
c = HookGroup()
|
||||||
|
for hook in self.hooks:
|
||||||
|
c.add(hook.clone())
|
||||||
|
return c
|
||||||
|
|
||||||
|
def clone_and_combine(self, other: 'HookGroup'):
|
||||||
|
c = self.clone()
|
||||||
|
if other is not None:
|
||||||
|
for hook in other.hooks:
|
||||||
|
c.add(hook.clone())
|
||||||
|
return c
|
||||||
|
|
||||||
|
def set_keyframes_on_hooks(self, hook_kf: 'HookKeyframeGroup'):
|
||||||
|
if hook_kf is None:
|
||||||
|
hook_kf = HookKeyframeGroup()
|
||||||
|
else:
|
||||||
|
hook_kf = hook_kf.clone()
|
||||||
|
for hook in self.hooks:
|
||||||
|
hook.hook_keyframe = hook_kf
|
||||||
|
|
||||||
|
def get_dict_repr(self):
|
||||||
|
d: dict[EnumHookType, dict[Hook, None]] = {}
|
||||||
|
for hook in self.hooks:
|
||||||
|
with_type = d.setdefault(hook.hook_type, {})
|
||||||
|
with_type[hook] = None
|
||||||
|
return d
|
||||||
|
|
||||||
|
def get_hooks_for_clip_schedule(self):
|
||||||
|
scheduled_hooks: dict[WeightHook, list[tuple[tuple[float,float], HookKeyframe]]] = {}
|
||||||
|
for hook in self.hooks:
|
||||||
|
# only care about WeightHooks, for now
|
||||||
|
if hook.hook_type == EnumHookType.Weight:
|
||||||
|
hook_schedule = []
|
||||||
|
# if no hook keyframes, assign default value
|
||||||
|
if len(hook.hook_keyframe.keyframes) == 0:
|
||||||
|
hook_schedule.append(((0.0, 1.0), None))
|
||||||
|
scheduled_hooks[hook] = hook_schedule
|
||||||
|
continue
|
||||||
|
# find ranges of values
|
||||||
|
prev_keyframe = hook.hook_keyframe.keyframes[0]
|
||||||
|
for keyframe in hook.hook_keyframe.keyframes:
|
||||||
|
if keyframe.start_percent > prev_keyframe.start_percent and not math.isclose(keyframe.strength, prev_keyframe.strength):
|
||||||
|
hook_schedule.append(((prev_keyframe.start_percent, keyframe.start_percent), prev_keyframe))
|
||||||
|
prev_keyframe = keyframe
|
||||||
|
elif keyframe.start_percent == prev_keyframe.start_percent:
|
||||||
|
prev_keyframe = keyframe
|
||||||
|
# create final range, assuming last start_percent was not 1.0
|
||||||
|
if not math.isclose(prev_keyframe.start_percent, 1.0):
|
||||||
|
hook_schedule.append(((prev_keyframe.start_percent, 1.0), prev_keyframe))
|
||||||
|
scheduled_hooks[hook] = hook_schedule
|
||||||
|
# hooks should not have their schedules in a list of tuples
|
||||||
|
all_ranges: list[tuple[float, float]] = []
|
||||||
|
for range_kfs in scheduled_hooks.values():
|
||||||
|
for t_range, keyframe in range_kfs:
|
||||||
|
all_ranges.append(t_range)
|
||||||
|
# turn list of ranges into boundaries
|
||||||
|
boundaries_set = set(itertools.chain.from_iterable(all_ranges))
|
||||||
|
boundaries_set.add(0.0)
|
||||||
|
boundaries = sorted(boundaries_set)
|
||||||
|
real_ranges = [(boundaries[i], boundaries[i + 1]) for i in range(len(boundaries) - 1)]
|
||||||
|
# with real ranges defined, give appropriate hooks w/ keyframes for each range
|
||||||
|
scheduled_keyframes: list[tuple[tuple[float,float], list[tuple[WeightHook, HookKeyframe]]]] = []
|
||||||
|
for t_range in real_ranges:
|
||||||
|
hooks_schedule = []
|
||||||
|
for hook, val in scheduled_hooks.items():
|
||||||
|
keyframe = None
|
||||||
|
# check if is a keyframe that works for the current t_range
|
||||||
|
for stored_range, stored_kf in val:
|
||||||
|
# if stored start is less than current end, then fits - give it assigned keyframe
|
||||||
|
if stored_range[0] < t_range[1] and stored_range[1] > t_range[0]:
|
||||||
|
keyframe = stored_kf
|
||||||
|
break
|
||||||
|
hooks_schedule.append((hook, keyframe))
|
||||||
|
scheduled_keyframes.append((t_range, hooks_schedule))
|
||||||
|
return scheduled_keyframes
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
for hook in self.hooks:
|
||||||
|
hook.reset()
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def combine_all_hooks(hooks_list: list['HookGroup'], require_count=0) -> 'HookGroup':
|
||||||
|
actual: list[HookGroup] = []
|
||||||
|
for group in hooks_list:
|
||||||
|
if group is not None:
|
||||||
|
actual.append(group)
|
||||||
|
if len(actual) < require_count:
|
||||||
|
raise Exception(f"Need at least {require_count} hooks to combine, but only had {len(actual)}.")
|
||||||
|
# if no hooks, then return None
|
||||||
|
if len(actual) == 0:
|
||||||
|
return None
|
||||||
|
# if only 1 hook, just return itself without cloning
|
||||||
|
elif len(actual) == 1:
|
||||||
|
return actual[0]
|
||||||
|
final_hook: HookGroup = None
|
||||||
|
for hook in actual:
|
||||||
|
if final_hook is None:
|
||||||
|
final_hook = hook.clone()
|
||||||
|
else:
|
||||||
|
final_hook = final_hook.clone_and_combine(hook)
|
||||||
|
return final_hook
|
||||||
|
|
||||||
|
|
||||||
|
class HookKeyframe:
|
||||||
|
def __init__(self, strength: float, start_percent=0.0, guarantee_steps=1):
|
||||||
|
self.strength = strength
|
||||||
|
# scheduling
|
||||||
|
self.start_percent = float(start_percent)
|
||||||
|
self.start_t = 999999999.9
|
||||||
|
self.guarantee_steps = guarantee_steps
|
||||||
|
|
||||||
|
def clone(self):
|
||||||
|
c = HookKeyframe(strength=self.strength,
|
||||||
|
start_percent=self.start_percent, guarantee_steps=self.guarantee_steps)
|
||||||
|
c.start_t = self.start_t
|
||||||
|
return c
|
||||||
|
|
||||||
|
class HookKeyframeGroup:
|
||||||
|
def __init__(self):
|
||||||
|
self.keyframes: list[HookKeyframe] = []
|
||||||
|
self._current_keyframe: HookKeyframe = None
|
||||||
|
self._current_used_steps = 0
|
||||||
|
self._current_index = 0
|
||||||
|
self._current_strength = None
|
||||||
|
self._curr_t = -1.
|
||||||
|
|
||||||
|
# properties shadow those of HookWeightsKeyframe
|
||||||
|
@property
|
||||||
|
def strength(self):
|
||||||
|
if self._current_keyframe is not None:
|
||||||
|
return self._current_keyframe.strength
|
||||||
|
return 1.0
|
||||||
|
|
||||||
|
def reset(self):
|
||||||
|
self._current_keyframe = None
|
||||||
|
self._current_used_steps = 0
|
||||||
|
self._current_index = 0
|
||||||
|
self._current_strength = None
|
||||||
|
self.curr_t = -1.
|
||||||
|
self._set_first_as_current()
|
||||||
|
|
||||||
|
def add(self, keyframe: HookKeyframe):
|
||||||
|
# add to end of list, then sort
|
||||||
|
self.keyframes.append(keyframe)
|
||||||
|
self.keyframes = get_sorted_list_via_attr(self.keyframes, "start_percent")
|
||||||
|
self._set_first_as_current()
|
||||||
|
|
||||||
|
def _set_first_as_current(self):
|
||||||
|
if len(self.keyframes) > 0:
|
||||||
|
self._current_keyframe = self.keyframes[0]
|
||||||
|
else:
|
||||||
|
self._current_keyframe = None
|
||||||
|
|
||||||
|
def has_index(self, index: int):
|
||||||
|
return index >= 0 and index < len(self.keyframes)
|
||||||
|
|
||||||
|
def is_empty(self):
|
||||||
|
return len(self.keyframes) == 0
|
||||||
|
|
||||||
|
def clone(self):
|
||||||
|
c = HookKeyframeGroup()
|
||||||
|
for keyframe in self.keyframes:
|
||||||
|
c.keyframes.append(keyframe.clone())
|
||||||
|
c._set_first_as_current()
|
||||||
|
return c
|
||||||
|
|
||||||
|
def initialize_timesteps(self, model: 'BaseModel'):
|
||||||
|
for keyframe in self.keyframes:
|
||||||
|
keyframe.start_t = model.model_sampling.percent_to_sigma(keyframe.start_percent)
|
||||||
|
|
||||||
|
def prepare_current_keyframe(self, curr_t: float) -> bool:
|
||||||
|
if self.is_empty():
|
||||||
|
return False
|
||||||
|
if curr_t == self._curr_t:
|
||||||
|
return False
|
||||||
|
prev_index = self._current_index
|
||||||
|
prev_strength = self._current_strength
|
||||||
|
# if met guaranteed steps, look for next keyframe in case need to switch
|
||||||
|
if self._current_used_steps >= self._current_keyframe.guarantee_steps:
|
||||||
|
# if has next index, loop through and see if need to switch
|
||||||
|
if self.has_index(self._current_index+1):
|
||||||
|
for i in range(self._current_index+1, len(self.keyframes)):
|
||||||
|
eval_c = self.keyframes[i]
|
||||||
|
# check if start_t is greater or equal to curr_t
|
||||||
|
# NOTE: t is in terms of sigmas, not percent, so bigger number = earlier step in sampling
|
||||||
|
if eval_c.start_t >= curr_t:
|
||||||
|
self._current_index = i
|
||||||
|
self._current_strength = eval_c.strength
|
||||||
|
self._current_keyframe = eval_c
|
||||||
|
self._current_used_steps = 0
|
||||||
|
# if guarantee_steps greater than zero, stop searching for other keyframes
|
||||||
|
if self._current_keyframe.guarantee_steps > 0:
|
||||||
|
break
|
||||||
|
# if eval_c is outside the percent range, stop looking further
|
||||||
|
else: break
|
||||||
|
# update steps current context is used
|
||||||
|
self._current_used_steps += 1
|
||||||
|
# update current timestep this was performed on
|
||||||
|
self._curr_t = curr_t
|
||||||
|
# return True if keyframe changed, False if no change
|
||||||
|
return prev_index != self._current_index and prev_strength != self._current_strength
|
||||||
|
|
||||||
|
|
||||||
|
class InterpolationMethod:
|
||||||
|
LINEAR = "linear"
|
||||||
|
EASE_IN = "ease_in"
|
||||||
|
EASE_OUT = "ease_out"
|
||||||
|
EASE_IN_OUT = "ease_in_out"
|
||||||
|
|
||||||
|
_LIST = [LINEAR, EASE_IN, EASE_OUT, EASE_IN_OUT]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_weights(cls, num_from: float, num_to: float, length: int, method: str, reverse=False):
|
||||||
|
diff = num_to - num_from
|
||||||
|
if method == cls.LINEAR:
|
||||||
|
weights = torch.linspace(num_from, num_to, length)
|
||||||
|
elif method == cls.EASE_IN:
|
||||||
|
index = torch.linspace(0, 1, length)
|
||||||
|
weights = diff * np.power(index, 2) + num_from
|
||||||
|
elif method == cls.EASE_OUT:
|
||||||
|
index = torch.linspace(0, 1, length)
|
||||||
|
weights = diff * (1 - np.power(1 - index, 2)) + num_from
|
||||||
|
elif method == cls.EASE_IN_OUT:
|
||||||
|
index = torch.linspace(0, 1, length)
|
||||||
|
weights = diff * ((1 - np.cos(index * np.pi)) / 2) + num_from
|
||||||
|
else:
|
||||||
|
raise ValueError(f"Unrecognized interpolation method '{method}'.")
|
||||||
|
if reverse:
|
||||||
|
weights = weights.flip(dims=(0,))
|
||||||
|
return weights
|
||||||
|
|
||||||
|
def get_sorted_list_via_attr(objects: list, attr: str) -> list:
|
||||||
|
if not objects:
|
||||||
|
return objects
|
||||||
|
elif len(objects) <= 1:
|
||||||
|
return [x for x in objects]
|
||||||
|
# now that we know we have to sort, do it following these rules:
|
||||||
|
# a) if objects have same value of attribute, maintain their relative order
|
||||||
|
# b) perform sorting of the groups of objects with same attributes
|
||||||
|
unique_attrs = {}
|
||||||
|
for o in objects:
|
||||||
|
val_attr = getattr(o, attr)
|
||||||
|
attr_list: list = unique_attrs.get(val_attr, list())
|
||||||
|
attr_list.append(o)
|
||||||
|
if val_attr not in unique_attrs:
|
||||||
|
unique_attrs[val_attr] = attr_list
|
||||||
|
# now that we have the unique attr values grouped together in relative order, sort them by key
|
||||||
|
sorted_attrs = dict(sorted(unique_attrs.items()))
|
||||||
|
# now flatten out the dict into a list to return
|
||||||
|
sorted_list = []
|
||||||
|
for object_list in sorted_attrs.values():
|
||||||
|
sorted_list.extend(object_list)
|
||||||
|
return sorted_list
|
||||||
|
|
||||||
|
def create_hook_lora(lora: dict[str, torch.Tensor], strength_model: float, strength_clip: float):
|
||||||
|
hook_group = HookGroup()
|
||||||
|
hook = WeightHook(strength_model=strength_model, strength_clip=strength_clip)
|
||||||
|
hook_group.add(hook)
|
||||||
|
hook.weights = lora
|
||||||
|
return hook_group
|
||||||
|
|
||||||
|
def create_hook_model_as_lora(weights_model, weights_clip, strength_model: float, strength_clip: float):
|
||||||
|
hook_group = HookGroup()
|
||||||
|
hook = WeightHook(strength_model=strength_model, strength_clip=strength_clip)
|
||||||
|
hook_group.add(hook)
|
||||||
|
patches_model = None
|
||||||
|
patches_clip = None
|
||||||
|
if weights_model is not None:
|
||||||
|
patches_model = {}
|
||||||
|
for key in weights_model:
|
||||||
|
patches_model[key] = ("model_as_lora", (weights_model[key],))
|
||||||
|
if weights_clip is not None:
|
||||||
|
patches_clip = {}
|
||||||
|
for key in weights_clip:
|
||||||
|
patches_clip[key] = ("model_as_lora", (weights_clip[key],))
|
||||||
|
hook.weights = patches_model
|
||||||
|
hook.weights_clip = patches_clip
|
||||||
|
hook.need_weight_init = False
|
||||||
|
return hook_group
|
||||||
|
|
||||||
|
def get_patch_weights_from_model(model: 'ModelPatcher', discard_model_sampling=True):
|
||||||
|
if model is None:
|
||||||
|
return None
|
||||||
|
patches_model: dict[str, torch.Tensor] = model.model.state_dict()
|
||||||
|
if discard_model_sampling:
|
||||||
|
# do not include ANY model_sampling components of the model that should act as a patch
|
||||||
|
for key in list(patches_model.keys()):
|
||||||
|
if key.startswith("model_sampling"):
|
||||||
|
patches_model.pop(key, None)
|
||||||
|
return patches_model
|
||||||
|
|
||||||
|
# NOTE: this function shows how to register weight hooks directly on the ModelPatchers
|
||||||
|
def load_hook_lora_for_models(model: 'ModelPatcher', clip: 'CLIP', lora: dict[str, torch.Tensor],
|
||||||
|
strength_model: float, strength_clip: float):
|
||||||
|
key_map = {}
|
||||||
|
if model is not None:
|
||||||
|
key_map = comfy.lora.model_lora_keys_unet(model.model, key_map)
|
||||||
|
if clip is not None:
|
||||||
|
key_map = comfy.lora.model_lora_keys_clip(clip.cond_stage_model, key_map)
|
||||||
|
|
||||||
|
hook_group = HookGroup()
|
||||||
|
hook = WeightHook()
|
||||||
|
hook_group.add(hook)
|
||||||
|
loaded: dict[str] = comfy.lora.load_lora(lora, key_map)
|
||||||
|
if model is not None:
|
||||||
|
new_modelpatcher = model.clone()
|
||||||
|
k = new_modelpatcher.add_hook_patches(hook=hook, patches=loaded, strength_patch=strength_model)
|
||||||
|
else:
|
||||||
|
k = ()
|
||||||
|
new_modelpatcher = None
|
||||||
|
|
||||||
|
if clip is not None:
|
||||||
|
new_clip = clip.clone()
|
||||||
|
k1 = new_clip.patcher.add_hook_patches(hook=hook, patches=loaded, strength_patch=strength_clip)
|
||||||
|
else:
|
||||||
|
k1 = ()
|
||||||
|
new_clip = None
|
||||||
|
k = set(k)
|
||||||
|
k1 = set(k1)
|
||||||
|
for x in loaded:
|
||||||
|
if (x not in k) and (x not in k1):
|
||||||
|
print(f"NOT LOADED {x}")
|
||||||
|
return (new_modelpatcher, new_clip, hook_group)
|
||||||
|
|
||||||
|
def _combine_hooks_from_values(c_dict: dict[str, HookGroup], values: dict[str, HookGroup], cache: dict[tuple[HookGroup, HookGroup], HookGroup]):
|
||||||
|
hooks_key = 'hooks'
|
||||||
|
# if hooks only exist in one dict, do what's needed so that it ends up in c_dict
|
||||||
|
if hooks_key not in values:
|
||||||
|
return
|
||||||
|
if hooks_key not in c_dict:
|
||||||
|
hooks_value = values.get(hooks_key, None)
|
||||||
|
if hooks_value is not None:
|
||||||
|
c_dict[hooks_key] = hooks_value
|
||||||
|
return
|
||||||
|
# otherwise, need to combine with minimum duplication via cache
|
||||||
|
hooks_tuple = (c_dict[hooks_key], values[hooks_key])
|
||||||
|
cached_hooks = cache.get(hooks_tuple, None)
|
||||||
|
if cached_hooks is None:
|
||||||
|
new_hooks = hooks_tuple[0].clone_and_combine(hooks_tuple[1])
|
||||||
|
cache[hooks_tuple] = new_hooks
|
||||||
|
c_dict[hooks_key] = new_hooks
|
||||||
|
else:
|
||||||
|
c_dict[hooks_key] = cache[hooks_tuple]
|
||||||
|
|
||||||
|
def conditioning_set_values_with_hooks(conditioning, values={}, append_hooks=True):
|
||||||
|
c = []
|
||||||
|
hooks_combine_cache: dict[tuple[HookGroup, HookGroup], HookGroup] = {}
|
||||||
|
for t in conditioning:
|
||||||
|
n = [t[0], t[1].copy()]
|
||||||
|
for k in values:
|
||||||
|
if append_hooks and k == 'hooks':
|
||||||
|
_combine_hooks_from_values(n[1], values, hooks_combine_cache)
|
||||||
|
else:
|
||||||
|
n[1][k] = values[k]
|
||||||
|
c.append(n)
|
||||||
|
|
||||||
|
return c
|
||||||
|
|
||||||
|
def set_hooks_for_conditioning(cond, hooks: HookGroup, append_hooks=True):
|
||||||
|
if hooks is None:
|
||||||
|
return cond
|
||||||
|
return conditioning_set_values_with_hooks(cond, {'hooks': hooks}, append_hooks=append_hooks)
|
||||||
|
|
||||||
|
def set_timesteps_for_conditioning(cond, timestep_range: tuple[float,float]):
|
||||||
|
if timestep_range is None:
|
||||||
|
return cond
|
||||||
|
return conditioning_set_values(cond, {"start_percent": timestep_range[0],
|
||||||
|
"end_percent": timestep_range[1]})
|
||||||
|
|
||||||
|
def set_mask_for_conditioning(cond, mask: torch.Tensor, set_cond_area: str, strength: float):
|
||||||
|
if mask is None:
|
||||||
|
return cond
|
||||||
|
set_area_to_bounds = False
|
||||||
|
if set_cond_area != 'default':
|
||||||
|
set_area_to_bounds = True
|
||||||
|
if len(mask.shape) < 3:
|
||||||
|
mask = mask.unsqueeze(0)
|
||||||
|
return conditioning_set_values(cond, {'mask': mask,
|
||||||
|
'set_area_to_bounds': set_area_to_bounds,
|
||||||
|
'mask_strength': strength})
|
||||||
|
|
||||||
|
def combine_conditioning(conds: list):
|
||||||
|
combined_conds = []
|
||||||
|
for cond in conds:
|
||||||
|
combined_conds.extend(cond)
|
||||||
|
return combined_conds
|
||||||
|
|
||||||
|
def combine_with_new_conds(conds: list, new_conds: list):
|
||||||
|
combined_conds = []
|
||||||
|
for c, new_c in zip(conds, new_conds):
|
||||||
|
combined_conds.append(combine_conditioning([c, new_c]))
|
||||||
|
return combined_conds
|
||||||
|
|
||||||
|
def set_conds_props(conds: list, strength: float, set_cond_area: str,
|
||||||
|
mask: torch.Tensor=None, hooks: HookGroup=None, timesteps_range: tuple[float,float]=None, append_hooks=True):
|
||||||
|
final_conds = []
|
||||||
|
for c in conds:
|
||||||
|
# first, apply lora_hook to conditioning, if provided
|
||||||
|
c = set_hooks_for_conditioning(c, hooks, append_hooks=append_hooks)
|
||||||
|
# next, apply mask to conditioning
|
||||||
|
c = set_mask_for_conditioning(cond=c, mask=mask, strength=strength, set_cond_area=set_cond_area)
|
||||||
|
# apply timesteps, if present
|
||||||
|
c = set_timesteps_for_conditioning(cond=c, timestep_range=timesteps_range)
|
||||||
|
# finally, apply mask to conditioning and store
|
||||||
|
final_conds.append(c)
|
||||||
|
return final_conds
|
||||||
|
|
||||||
|
def set_conds_props_and_combine(conds: list, new_conds: list, strength: float=1.0, set_cond_area: str="default",
|
||||||
|
mask: torch.Tensor=None, hooks: HookGroup=None, timesteps_range: tuple[float,float]=None, append_hooks=True):
|
||||||
|
combined_conds = []
|
||||||
|
for c, masked_c in zip(conds, new_conds):
|
||||||
|
# first, apply lora_hook to new conditioning, if provided
|
||||||
|
masked_c = set_hooks_for_conditioning(masked_c, hooks, append_hooks=append_hooks)
|
||||||
|
# next, apply mask to new conditioning, if provided
|
||||||
|
masked_c = set_mask_for_conditioning(cond=masked_c, mask=mask, set_cond_area=set_cond_area, strength=strength)
|
||||||
|
# apply timesteps, if present
|
||||||
|
masked_c = set_timesteps_for_conditioning(cond=masked_c, timestep_range=timesteps_range)
|
||||||
|
# finally, combine with existing conditioning and store
|
||||||
|
combined_conds.append(combine_conditioning([c, masked_c]))
|
||||||
|
return combined_conds
|
||||||
|
|
||||||
|
def set_default_conds_and_combine(conds: list, new_conds: list,
|
||||||
|
hooks: HookGroup=None, timesteps_range: tuple[float,float]=None, append_hooks=True):
|
||||||
|
combined_conds = []
|
||||||
|
for c, new_c in zip(conds, new_conds):
|
||||||
|
# first, apply lora_hook to new conditioning, if provided
|
||||||
|
new_c = set_hooks_for_conditioning(new_c, hooks, append_hooks=append_hooks)
|
||||||
|
# next, add default_cond key to cond so that during sampling, it can be identified
|
||||||
|
new_c = conditioning_set_values(new_c, {'default': True})
|
||||||
|
# apply timesteps, if present
|
||||||
|
new_c = set_timesteps_for_conditioning(cond=new_c, timestep_range=timesteps_range)
|
||||||
|
# finally, combine with existing conditioning and store
|
||||||
|
combined_conds.append(combine_conditioning([c, new_c]))
|
||||||
|
return combined_conds
|
||||||
@ -175,12 +175,14 @@ def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, dis
|
|||||||
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
||||||
if callback is not None:
|
if callback is not None:
|
||||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||||
d = to_d(x, sigmas[i], denoised)
|
|
||||||
# Euler method
|
if sigma_down == 0:
|
||||||
dt = sigma_down - sigmas[i]
|
x = denoised
|
||||||
x = x + d * dt
|
else:
|
||||||
if sigmas[i + 1] > 0:
|
d = to_d(x, sigmas[i], denoised)
|
||||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
|
# Euler method
|
||||||
|
dt = sigma_down - sigmas[i]
|
||||||
|
x = x + d * dt + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
|
||||||
return x
|
return x
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
@ -192,19 +194,22 @@ def sample_euler_ancestral_RF(model, x, sigmas, extra_args=None, callback=None,
|
|||||||
for i in trange(len(sigmas) - 1, disable=disable):
|
for i in trange(len(sigmas) - 1, disable=disable):
|
||||||
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||||
# sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
# sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)
|
||||||
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
|
|
||||||
sigma_down = sigmas[i+1] * downstep_ratio
|
|
||||||
alpha_ip1 = 1 - sigmas[i+1]
|
|
||||||
alpha_down = 1 - sigma_down
|
|
||||||
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
|
|
||||||
if callback is not None:
|
if callback is not None:
|
||||||
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||||
|
|
||||||
# Euler method
|
if sigmas[i + 1] == 0:
|
||||||
sigma_down_i_ratio = sigma_down / sigmas[i]
|
x = denoised
|
||||||
x = sigma_down_i_ratio * x + (1 - sigma_down_i_ratio) * denoised
|
else:
|
||||||
if sigmas[i + 1] > 0 and eta > 0:
|
downstep_ratio = 1 + (sigmas[i + 1] / sigmas[i] - 1) * eta
|
||||||
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
|
sigma_down = sigmas[i + 1] * downstep_ratio
|
||||||
|
alpha_ip1 = 1 - sigmas[i + 1]
|
||||||
|
alpha_down = 1 - sigma_down
|
||||||
|
renoise_coeff = (sigmas[i + 1]**2 - sigma_down**2 * alpha_ip1**2 / alpha_down**2)**0.5
|
||||||
|
# Euler method
|
||||||
|
sigma_down_i_ratio = sigma_down / sigmas[i]
|
||||||
|
x = sigma_down_i_ratio * x + (1 - sigma_down_i_ratio) * denoised
|
||||||
|
if eta > 0:
|
||||||
|
x = (alpha_ip1 / alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
|
||||||
return x
|
return x
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
@ -280,6 +285,9 @@ def sample_dpm_2(model, x, sigmas, extra_args=None, callback=None, disable=None,
|
|||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
def sample_dpm_2_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
def sample_dpm_2_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||||
|
if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):
|
||||||
|
return sample_dpm_2_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler)
|
||||||
|
|
||||||
"""Ancestral sampling with DPM-Solver second-order steps."""
|
"""Ancestral sampling with DPM-Solver second-order steps."""
|
||||||
extra_args = {} if extra_args is None else extra_args
|
extra_args = {} if extra_args is None else extra_args
|
||||||
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||||
@ -306,6 +314,38 @@ def sample_dpm_2_ancestral(model, x, sigmas, extra_args=None, callback=None, dis
|
|||||||
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
|
x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
@torch.no_grad()
|
||||||
|
def sample_dpm_2_ancestral_RF(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):
|
||||||
|
"""Ancestral sampling with DPM-Solver second-order steps."""
|
||||||
|
extra_args = {} if extra_args is None else extra_args
|
||||||
|
noise_sampler = default_noise_sampler(x) if noise_sampler is None else noise_sampler
|
||||||
|
s_in = x.new_ones([x.shape[0]])
|
||||||
|
for i in trange(len(sigmas) - 1, disable=disable):
|
||||||
|
denoised = model(x, sigmas[i] * s_in, **extra_args)
|
||||||
|
downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta
|
||||||
|
sigma_down = sigmas[i+1] * downstep_ratio
|
||||||
|
alpha_ip1 = 1 - sigmas[i+1]
|
||||||
|
alpha_down = 1 - sigma_down
|
||||||
|
renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5
|
||||||
|
|
||||||
|
if callback is not None:
|
||||||
|
callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})
|
||||||
|
d = to_d(x, sigmas[i], denoised)
|
||||||
|
if sigma_down == 0:
|
||||||
|
# Euler method
|
||||||
|
dt = sigma_down - sigmas[i]
|
||||||
|
x = x + d * dt
|
||||||
|
else:
|
||||||
|
# DPM-Solver-2
|
||||||
|
sigma_mid = sigmas[i].log().lerp(sigma_down.log(), 0.5).exp()
|
||||||
|
dt_1 = sigma_mid - sigmas[i]
|
||||||
|
dt_2 = sigma_down - sigmas[i]
|
||||||
|
x_2 = x + d * dt_1
|
||||||
|
denoised_2 = model(x_2, sigma_mid * s_in, **extra_args)
|
||||||
|
d_2 = to_d(x_2, sigma_mid, denoised_2)
|
||||||
|
x = x + d_2 * dt_2
|
||||||
|
x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff
|
||||||
|
return x
|
||||||
|
|
||||||
def linear_multistep_coeff(order, t, i, j):
|
def linear_multistep_coeff(order, t, i, j):
|
||||||
if order - 1 > i:
|
if order - 1 > i:
|
||||||
|
|||||||
@ -219,4 +219,136 @@ class Mochi(LatentFormat):
|
|||||||
|
|
||||||
class LTXV(LatentFormat):
|
class LTXV(LatentFormat):
|
||||||
latent_channels = 128
|
latent_channels = 128
|
||||||
|
def __init__(self):
|
||||||
|
self.latent_rgb_factors = [
|
||||||
|
[ 1.1202e-02, -6.3815e-04, -1.0021e-02],
|
||||||
|
[ 8.6031e-02, 6.5813e-02, 9.5409e-04],
|
||||||
|
[-1.2576e-02, -7.5734e-03, -4.0528e-03],
|
||||||
|
[ 9.4063e-03, -2.1688e-03, 2.6093e-03],
|
||||||
|
[ 3.7636e-03, 1.2765e-02, 9.1548e-03],
|
||||||
|
[ 2.1024e-02, -5.2973e-03, 3.4373e-03],
|
||||||
|
[-8.8896e-03, -1.9703e-02, -1.8761e-02],
|
||||||
|
[-1.3160e-02, -1.0523e-02, 1.9709e-03],
|
||||||
|
[-1.5152e-03, -6.9891e-03, -7.5810e-03],
|
||||||
|
[-1.7247e-03, 4.6560e-04, -3.3839e-03],
|
||||||
|
[ 1.3617e-02, 4.7077e-03, -2.0045e-03],
|
||||||
|
[ 1.0256e-02, 7.7318e-03, 1.3948e-02],
|
||||||
|
[-1.6108e-02, -6.2151e-03, 1.1561e-03],
|
||||||
|
[ 7.3407e-03, 1.5628e-02, 4.4865e-04],
|
||||||
|
[ 9.5357e-04, -2.9518e-03, -1.4760e-02],
|
||||||
|
[ 1.9143e-02, 1.0868e-02, 1.2264e-02],
|
||||||
|
[ 4.4575e-03, 3.6682e-05, -6.8508e-03],
|
||||||
|
[-4.5681e-04, 3.2570e-03, 7.7929e-03],
|
||||||
|
[ 3.3902e-02, 3.3405e-02, 3.7454e-02],
|
||||||
|
[-2.3001e-02, -2.4877e-03, -3.1033e-03],
|
||||||
|
[ 5.0265e-02, 3.8841e-02, 3.3539e-02],
|
||||||
|
[-4.1018e-03, -1.1095e-03, 1.5859e-03],
|
||||||
|
[-1.2689e-01, -1.3107e-01, -2.1005e-01],
|
||||||
|
[ 2.6276e-02, 1.4189e-02, -3.5963e-03],
|
||||||
|
[-4.8679e-03, 8.8486e-03, 7.8029e-03],
|
||||||
|
[-1.6610e-03, -4.8597e-03, -5.2060e-03],
|
||||||
|
[-2.1010e-03, 2.3610e-03, 9.3796e-03],
|
||||||
|
[-2.2482e-02, -2.1305e-02, -1.5087e-02],
|
||||||
|
[-1.5753e-02, -1.0646e-02, -6.5083e-03],
|
||||||
|
[-4.6975e-03, 5.0288e-03, -6.7390e-03],
|
||||||
|
[ 1.1951e-02, 2.0712e-02, 1.6191e-02],
|
||||||
|
[-6.3704e-03, -8.4827e-03, -9.5483e-03],
|
||||||
|
[ 7.2610e-03, -9.9326e-03, -2.2978e-02],
|
||||||
|
[-9.1904e-04, 6.2882e-03, 9.5720e-03],
|
||||||
|
[-3.7178e-02, -3.7123e-02, -5.6713e-02],
|
||||||
|
[-1.3373e-01, -1.0720e-01, -5.3801e-02],
|
||||||
|
[-5.3702e-03, 8.1256e-03, 8.8397e-03],
|
||||||
|
[-1.5247e-01, -2.1437e-01, -2.1843e-01],
|
||||||
|
[ 3.1441e-02, 7.0335e-03, -9.7541e-03],
|
||||||
|
[ 2.1528e-03, -8.9817e-03, -2.1023e-02],
|
||||||
|
[ 3.8461e-03, -5.8957e-03, -1.5014e-02],
|
||||||
|
[-4.3470e-03, -1.2940e-02, -1.5972e-02],
|
||||||
|
[-5.4781e-03, -1.0842e-02, -3.0204e-03],
|
||||||
|
[-6.5347e-03, 3.0806e-03, -1.0163e-02],
|
||||||
|
[-5.0414e-03, -7.1503e-03, -8.9686e-04],
|
||||||
|
[-8.5851e-03, -2.4351e-03, 1.0674e-03],
|
||||||
|
[-9.0016e-03, -9.6493e-03, 1.5692e-03],
|
||||||
|
[ 5.0914e-03, 1.2099e-02, 1.9968e-02],
|
||||||
|
[ 1.3758e-02, 1.1669e-02, 8.1958e-03],
|
||||||
|
[-1.0518e-02, -1.1575e-02, -4.1307e-03],
|
||||||
|
[-2.8410e-02, -3.1266e-02, -2.2149e-02],
|
||||||
|
[ 2.9336e-03, 3.6511e-02, 1.8717e-02],
|
||||||
|
[-1.6703e-02, -1.6696e-02, -4.4529e-03],
|
||||||
|
[ 4.8818e-02, 4.0063e-02, 8.7410e-03],
|
||||||
|
[-1.5066e-02, -5.7328e-04, 2.9785e-03],
|
||||||
|
[-1.7613e-02, -8.1034e-03, 1.3086e-02],
|
||||||
|
[-9.2633e-03, 1.0803e-02, -6.3489e-03],
|
||||||
|
[ 3.0851e-03, 4.7750e-04, 1.2347e-02],
|
||||||
|
[-2.2785e-02, -2.3043e-02, -2.6005e-02],
|
||||||
|
[-2.4787e-02, -1.5389e-02, -2.2104e-02],
|
||||||
|
[-2.3572e-02, 1.0544e-03, 1.2361e-02],
|
||||||
|
[-7.8915e-03, -1.2271e-03, -6.0968e-03],
|
||||||
|
[-1.1478e-02, -1.2543e-03, 6.2679e-03],
|
||||||
|
[-5.4229e-02, 2.6644e-02, 6.3394e-03],
|
||||||
|
[ 4.4216e-03, -7.3338e-03, -1.0464e-02],
|
||||||
|
[-4.5013e-03, 1.6082e-03, 1.4420e-02],
|
||||||
|
[ 1.3673e-02, 8.8877e-03, 4.1253e-03],
|
||||||
|
[-1.0145e-02, 9.0072e-03, 1.5695e-02],
|
||||||
|
[-5.6234e-03, 1.1847e-03, 8.1261e-03],
|
||||||
|
[-3.7171e-03, -5.3538e-03, 1.2590e-03],
|
||||||
|
[ 2.9476e-02, 2.1424e-02, 3.0424e-02],
|
||||||
|
[-3.4925e-02, -2.4340e-02, -2.5316e-02],
|
||||||
|
[-3.4127e-02, -2.2406e-02, -1.0589e-02],
|
||||||
|
[-1.7342e-02, -1.3249e-02, -1.0719e-02],
|
||||||
|
[-2.1478e-03, -8.6051e-03, -2.9878e-03],
|
||||||
|
[ 1.2089e-03, -4.2391e-03, -6.8569e-03],
|
||||||
|
[ 9.0411e-04, -6.6886e-03, -6.7547e-05],
|
||||||
|
[ 1.6048e-02, -1.0057e-02, -2.8929e-02],
|
||||||
|
[ 1.2290e-03, 1.0163e-02, 1.8861e-02],
|
||||||
|
[ 1.7264e-02, 2.7257e-04, 1.3785e-02],
|
||||||
|
[-1.3482e-02, -3.6427e-03, 6.7481e-04],
|
||||||
|
[ 4.6782e-03, -5.2423e-03, 2.4467e-03],
|
||||||
|
[-5.9113e-03, -6.2244e-03, -1.8162e-03],
|
||||||
|
[ 1.5496e-02, 1.4582e-02, 1.9514e-03],
|
||||||
|
[ 7.4958e-03, 1.5886e-03, -8.2305e-03],
|
||||||
|
[ 1.9086e-02, 1.6360e-03, -3.9674e-03],
|
||||||
|
[-5.7021e-03, -2.7307e-03, -4.1066e-03],
|
||||||
|
[ 1.7450e-03, 1.4602e-02, 2.5794e-02],
|
||||||
|
[-8.2788e-04, 2.2902e-03, 4.5161e-03],
|
||||||
|
[ 1.1632e-02, 8.9193e-03, -7.2813e-03],
|
||||||
|
[ 7.5721e-03, 2.6784e-03, 1.1393e-02],
|
||||||
|
[ 5.1939e-03, 3.6903e-03, 1.4049e-02],
|
||||||
|
[-1.8383e-02, -2.2529e-02, -2.4477e-02],
|
||||||
|
[ 5.8842e-04, -5.7874e-03, -1.4770e-02],
|
||||||
|
[-1.6125e-02, -8.6101e-03, -1.4533e-02],
|
||||||
|
[ 2.0540e-02, 2.0729e-02, 6.4338e-03],
|
||||||
|
[ 3.3587e-03, -1.1226e-02, -1.6444e-02],
|
||||||
|
[-1.4742e-03, -1.0489e-02, 1.7097e-03],
|
||||||
|
[ 2.8130e-02, 2.3546e-02, 3.2791e-02],
|
||||||
|
[-1.8532e-02, -1.2842e-02, -8.7756e-03],
|
||||||
|
[-8.0533e-03, -1.0771e-02, -1.7536e-02],
|
||||||
|
[-3.9009e-03, 1.6150e-02, 3.3359e-02],
|
||||||
|
[-7.4554e-03, -1.4154e-02, -6.1910e-03],
|
||||||
|
[ 3.4734e-03, -1.1370e-02, -1.0581e-02],
|
||||||
|
[ 1.1476e-02, 3.9281e-03, 2.8231e-03],
|
||||||
|
[ 7.1639e-03, -1.4741e-03, -3.8066e-03],
|
||||||
|
[ 2.2250e-03, -8.7552e-03, -9.5719e-03],
|
||||||
|
[ 2.4146e-02, 2.1696e-02, 2.8056e-02],
|
||||||
|
[-5.4365e-03, -2.4291e-02, -1.7802e-02],
|
||||||
|
[ 7.4263e-03, 1.0510e-02, 1.2705e-02],
|
||||||
|
[ 6.2669e-03, 6.2658e-03, 1.9211e-02],
|
||||||
|
[ 1.6378e-02, 9.4933e-03, 6.6971e-03],
|
||||||
|
[ 1.7173e-02, 2.3601e-02, 2.3296e-02],
|
||||||
|
[-1.4568e-02, -9.8279e-03, -1.1556e-02],
|
||||||
|
[ 1.4431e-02, 1.4430e-02, 6.6362e-03],
|
||||||
|
[-6.8230e-03, 1.8863e-02, 1.4555e-02],
|
||||||
|
[ 6.1156e-03, 3.4700e-03, -2.6662e-03],
|
||||||
|
[-2.6983e-03, -5.9402e-03, -9.2276e-03],
|
||||||
|
[ 1.0235e-02, 7.4173e-03, -7.6243e-03],
|
||||||
|
[-1.3255e-02, 1.9322e-02, -9.2153e-04],
|
||||||
|
[ 2.4222e-03, -4.8039e-03, -1.5759e-02],
|
||||||
|
[ 2.6244e-02, 2.5951e-02, 2.0249e-02],
|
||||||
|
[ 1.5711e-02, 1.8498e-02, 2.7407e-03],
|
||||||
|
[-2.1714e-03, 4.7214e-03, -2.2443e-02],
|
||||||
|
[-7.4747e-03, 7.4166e-03, 1.4430e-02],
|
||||||
|
[-8.3906e-03, -7.9776e-03, 9.7927e-03],
|
||||||
|
[ 3.8321e-02, 9.6622e-03, -1.9268e-02],
|
||||||
|
[-1.4605e-02, -6.7032e-03, 3.9675e-03]
|
||||||
|
]
|
||||||
|
|
||||||
|
self.latent_rgb_factors_bias = [-0.0571, -0.1657, -0.2512]
|
||||||
|
|||||||
@ -2,7 +2,7 @@
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from typing import Literal, Dict, Any
|
from typing import Literal
|
||||||
import math
|
import math
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
ops = comfy.ops.disable_weight_init
|
ops = comfy.ops.disable_weight_init
|
||||||
|
|||||||
@ -2,8 +2,8 @@
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
from torch import Tensor, einsum
|
from torch import Tensor
|
||||||
from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, TypeVar, Union
|
from typing import List, Union
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
import math
|
import math
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
|
|||||||
@ -16,7 +16,6 @@
|
|||||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import torch
|
|
||||||
import torchvision
|
import torchvision
|
||||||
from torch import nn
|
from torch import nn
|
||||||
from .common import LayerNorm2d_op
|
from .common import LayerNorm2d_op
|
||||||
|
|||||||
@ -2,7 +2,7 @@ import torch
|
|||||||
import comfy.ops
|
import comfy.ops
|
||||||
|
|
||||||
def pad_to_patch_size(img, patch_size=(2, 2), padding_mode="circular"):
|
def pad_to_patch_size(img, patch_size=(2, 2), padding_mode="circular"):
|
||||||
if padding_mode == "circular" and torch.jit.is_tracing() or torch.jit.is_scripting():
|
if padding_mode == "circular" and (torch.jit.is_tracing() or torch.jit.is_scripting()):
|
||||||
padding_mode = "reflect"
|
padding_mode = "reflect"
|
||||||
pad_h = (patch_size[0] - img.shape[-2] % patch_size[0]) % patch_size[0]
|
pad_h = (patch_size[0] - img.shape[-2] % patch_size[0]) % patch_size[0]
|
||||||
pad_w = (patch_size[1] - img.shape[-1] % patch_size[1]) % patch_size[1]
|
pad_w = (patch_size[1] - img.shape[-1] % patch_size[1]) % patch_size[1]
|
||||||
|
|||||||
@ -6,9 +6,7 @@ import math
|
|||||||
from torch import Tensor, nn
|
from torch import Tensor, nn
|
||||||
from einops import rearrange, repeat
|
from einops import rearrange, repeat
|
||||||
|
|
||||||
from .layers import (DoubleStreamBlock, EmbedND, LastLayer,
|
from .layers import (timestep_embedding)
|
||||||
MLPEmbedder, SingleStreamBlock,
|
|
||||||
timestep_embedding)
|
|
||||||
|
|
||||||
from .model import Flux
|
from .model import Flux
|
||||||
import comfy.ldm.common_dit
|
import comfy.ldm.common_dit
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
#original code from https://github.com/genmoai/models under apache 2.0 license
|
#original code from https://github.com/genmoai/models under apache 2.0 license
|
||||||
#adapted to ComfyUI
|
#adapted to ComfyUI
|
||||||
|
|
||||||
from typing import Optional, Tuple
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|||||||
@ -1,7 +1,7 @@
|
|||||||
#original code from https://github.com/genmoai/models under apache 2.0 license
|
#original code from https://github.com/genmoai/models under apache 2.0 license
|
||||||
#adapted to ComfyUI
|
#adapted to ComfyUI
|
||||||
|
|
||||||
from typing import Callable, List, Optional, Tuple, Union
|
from typing import List, Optional, Tuple, Union
|
||||||
from functools import partial
|
from functools import partial
|
||||||
import math
|
import math
|
||||||
|
|
||||||
|
|||||||
@ -1,24 +1,17 @@
|
|||||||
from typing import Any, Optional
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
from torch.utils import checkpoint
|
|
||||||
|
|
||||||
from comfy.ldm.modules.diffusionmodules.mmdit import (
|
from comfy.ldm.modules.diffusionmodules.mmdit import (
|
||||||
Mlp,
|
|
||||||
TimestepEmbedder,
|
TimestepEmbedder,
|
||||||
PatchEmbed,
|
PatchEmbed,
|
||||||
RMSNorm,
|
|
||||||
)
|
)
|
||||||
from comfy.ldm.modules.diffusionmodules.util import timestep_embedding
|
|
||||||
from .poolers import AttentionPool
|
from .poolers import AttentionPool
|
||||||
|
|
||||||
import comfy.latent_formats
|
import comfy.latent_formats
|
||||||
from .models import HunYuanDiTBlock, calc_rope
|
from .models import HunYuanDiTBlock, calc_rope
|
||||||
|
|
||||||
from .posemb_layers import get_2d_rotary_pos_embed, get_fill_resize_and_crop
|
|
||||||
|
|
||||||
|
|
||||||
class HunYuanControlNet(nn.Module):
|
class HunYuanControlNet(nn.Module):
|
||||||
|
|||||||
@ -1,8 +1,6 @@
|
|||||||
from typing import Any
|
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
|
||||||
|
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
from comfy.ldm.modules.diffusionmodules.mmdit import Mlp, TimestepEmbedder, PatchEmbed, RMSNorm
|
from comfy.ldm.modules.diffusionmodules.mmdit import Mlp, TimestepEmbedder, PatchEmbed, RMSNorm
|
||||||
|
|||||||
@ -1,6 +1,5 @@
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import torch.nn.functional as F
|
|
||||||
from comfy.ldm.modules.attention import optimized_attention
|
from comfy.ldm.modules.attention import optimized_attention
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
|
|
||||||
|
|||||||
@ -379,6 +379,7 @@ class LTXVModel(torch.nn.Module):
|
|||||||
positional_embedding_max_pos=[20, 2048, 2048],
|
positional_embedding_max_pos=[20, 2048, 2048],
|
||||||
dtype=None, device=None, operations=None, **kwargs):
|
dtype=None, device=None, operations=None, **kwargs):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
self.generator = None
|
||||||
self.dtype = dtype
|
self.dtype = dtype
|
||||||
self.out_channels = in_channels
|
self.out_channels = in_channels
|
||||||
self.inner_dim = num_attention_heads * attention_head_dim
|
self.inner_dim = num_attention_heads * attention_head_dim
|
||||||
@ -415,7 +416,7 @@ class LTXVModel(torch.nn.Module):
|
|||||||
|
|
||||||
self.patchifier = SymmetricPatchifier(1)
|
self.patchifier = SymmetricPatchifier(1)
|
||||||
|
|
||||||
def forward(self, x, timestep, context, attention_mask, frame_rate=25, guiding_latent=None, transformer_options={}, **kwargs):
|
def forward(self, x, timestep, context, attention_mask, frame_rate=25, guiding_latent=None, guiding_latent_noise_scale=0, transformer_options={}, **kwargs):
|
||||||
patches_replace = transformer_options.get("patches_replace", {})
|
patches_replace = transformer_options.get("patches_replace", {})
|
||||||
|
|
||||||
indices_grid = self.patchifier.get_grid(
|
indices_grid = self.patchifier.get_grid(
|
||||||
@ -431,10 +432,22 @@ class LTXVModel(torch.nn.Module):
|
|||||||
ts = torch.ones([x.shape[0], 1, x.shape[2], x.shape[3], x.shape[4]], device=x.device, dtype=x.dtype)
|
ts = torch.ones([x.shape[0], 1, x.shape[2], x.shape[3], x.shape[4]], device=x.device, dtype=x.dtype)
|
||||||
input_ts = timestep.view([timestep.shape[0]] + [1] * (x.ndim - 1))
|
input_ts = timestep.view([timestep.shape[0]] + [1] * (x.ndim - 1))
|
||||||
ts *= input_ts
|
ts *= input_ts
|
||||||
ts[:, :, 0] = 0.0
|
ts[:, :, 0] = guiding_latent_noise_scale * (input_ts[:, :, 0] ** 2)
|
||||||
timestep = self.patchifier.patchify(ts)
|
timestep = self.patchifier.patchify(ts)
|
||||||
input_x = x.clone()
|
input_x = x.clone()
|
||||||
x[:, :, 0] = guiding_latent[:, :, 0]
|
x[:, :, 0] = guiding_latent[:, :, 0]
|
||||||
|
if guiding_latent_noise_scale > 0:
|
||||||
|
if self.generator is None:
|
||||||
|
self.generator = torch.Generator(device=x.device).manual_seed(42)
|
||||||
|
elif self.generator.device != x.device:
|
||||||
|
self.generator = torch.Generator(device=x.device).set_state(self.generator.get_state())
|
||||||
|
|
||||||
|
noise_shape = [guiding_latent.shape[0], guiding_latent.shape[1], 1, guiding_latent.shape[3], guiding_latent.shape[4]]
|
||||||
|
scale = guiding_latent_noise_scale * (input_ts ** 2)
|
||||||
|
guiding_noise = scale * torch.randn(size=noise_shape, device=x.device, generator=self.generator)
|
||||||
|
|
||||||
|
x[:, :, 0] = guiding_noise[:, :, 0] + x[:, :, 0] * (1.0 - scale[:, :, 0])
|
||||||
|
|
||||||
|
|
||||||
orig_shape = list(x.shape)
|
orig_shape = list(x.shape)
|
||||||
|
|
||||||
|
|||||||
@ -3,7 +3,7 @@ from torch import nn
|
|||||||
from functools import partial
|
from functools import partial
|
||||||
import math
|
import math
|
||||||
from einops import rearrange
|
from einops import rearrange
|
||||||
from typing import Any, Mapping, Optional, Tuple, Union, List
|
from typing import Optional, Tuple, Union
|
||||||
from .conv_nd_factory import make_conv_nd, make_linear_nd
|
from .conv_nd_factory import make_conv_nd, make_linear_nd
|
||||||
from .pixel_norm import PixelNorm
|
from .pixel_norm import PixelNorm
|
||||||
|
|
||||||
|
|||||||
@ -1,6 +1,5 @@
|
|||||||
from typing import Tuple, Union
|
from typing import Tuple, Union
|
||||||
|
|
||||||
import torch
|
|
||||||
|
|
||||||
from .dual_conv3d import DualConv3d
|
from .dual_conv3d import DualConv3d
|
||||||
from .causal_conv3d import CausalConv3d
|
from .causal_conv3d import CausalConv3d
|
||||||
|
|||||||
@ -1,6 +1,6 @@
|
|||||||
import torch
|
import torch
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
from typing import Any, Dict, Tuple, Union
|
||||||
|
|
||||||
from comfy.ldm.modules.distributions.distributions import DiagonalGaussianDistribution
|
from comfy.ldm.modules.distributions.distributions import DiagonalGaussianDistribution
|
||||||
|
|
||||||
|
|||||||
@ -1,5 +1,3 @@
|
|||||||
import logging
|
|
||||||
import math
|
|
||||||
from typing import Dict, Optional, List
|
from typing import Dict, Optional, List
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
|
|||||||
@ -3,7 +3,6 @@ import math
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from typing import Optional, Any
|
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
|
|||||||
@ -9,12 +9,12 @@ import logging
|
|||||||
from .util import (
|
from .util import (
|
||||||
checkpoint,
|
checkpoint,
|
||||||
avg_pool_nd,
|
avg_pool_nd,
|
||||||
zero_module,
|
|
||||||
timestep_embedding,
|
timestep_embedding,
|
||||||
AlphaBlender,
|
AlphaBlender,
|
||||||
)
|
)
|
||||||
from ..attention import SpatialTransformer, SpatialVideoTransformer, default
|
from ..attention import SpatialTransformer, SpatialVideoTransformer, default
|
||||||
from comfy.ldm.util import exists
|
from comfy.ldm.util import exists
|
||||||
|
import comfy.patcher_extension
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
ops = comfy.ops.disable_weight_init
|
ops = comfy.ops.disable_weight_init
|
||||||
|
|
||||||
@ -47,6 +47,15 @@ def forward_timestep_embed(ts, x, emb, context=None, transformer_options={}, out
|
|||||||
elif isinstance(layer, Upsample):
|
elif isinstance(layer, Upsample):
|
||||||
x = layer(x, output_shape=output_shape)
|
x = layer(x, output_shape=output_shape)
|
||||||
else:
|
else:
|
||||||
|
if "patches" in transformer_options and "forward_timestep_embed_patch" in transformer_options["patches"]:
|
||||||
|
found_patched = False
|
||||||
|
for class_type, handler in transformer_options["patches"]["forward_timestep_embed_patch"]:
|
||||||
|
if isinstance(layer, class_type):
|
||||||
|
x = handler(layer, x, emb, context, transformer_options, output_shape, time_context, num_video_frames, image_only_indicator)
|
||||||
|
found_patched = True
|
||||||
|
break
|
||||||
|
if found_patched:
|
||||||
|
continue
|
||||||
x = layer(x)
|
x = layer(x)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
@ -819,6 +828,13 @@ class UNetModel(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
def forward(self, x, timesteps=None, context=None, y=None, control=None, transformer_options={}, **kwargs):
|
def forward(self, x, timesteps=None, context=None, y=None, control=None, transformer_options={}, **kwargs):
|
||||||
|
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||||
|
self._forward,
|
||||||
|
self,
|
||||||
|
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.DIFFUSION_MODEL, transformer_options)
|
||||||
|
).execute(x, timesteps, context, y, control, transformer_options, **kwargs)
|
||||||
|
|
||||||
|
def _forward(self, x, timesteps=None, context=None, y=None, control=None, transformer_options={}, **kwargs):
|
||||||
"""
|
"""
|
||||||
Apply the model to an input batch.
|
Apply the model to an input batch.
|
||||||
:param x: an [N x C x ...] Tensor of inputs.
|
:param x: an [N x C x ...] Tensor of inputs.
|
||||||
|
|||||||
@ -4,7 +4,6 @@ import numpy as np
|
|||||||
from functools import partial
|
from functools import partial
|
||||||
|
|
||||||
from .util import extract_into_tensor, make_beta_schedule
|
from .util import extract_into_tensor, make_beta_schedule
|
||||||
from comfy.ldm.util import default
|
|
||||||
|
|
||||||
|
|
||||||
class AbstractLowScaleModel(nn.Module):
|
class AbstractLowScaleModel(nn.Module):
|
||||||
|
|||||||
@ -8,7 +8,6 @@
|
|||||||
# thanks!
|
# thanks!
|
||||||
|
|
||||||
|
|
||||||
import os
|
|
||||||
import math
|
import math
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
|||||||
@ -1,5 +1,5 @@
|
|||||||
import functools
|
import functools
|
||||||
from typing import Callable, Iterable, Union
|
from typing import Iterable, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from einops import rearrange, repeat
|
from einops import rearrange, repeat
|
||||||
|
|||||||
@ -33,7 +33,7 @@ LORA_CLIP_MAP = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
def load_lora(lora, to_load):
|
def load_lora(lora, to_load, log_missing=True):
|
||||||
patch_dict = {}
|
patch_dict = {}
|
||||||
loaded_keys = set()
|
loaded_keys = set()
|
||||||
for x in to_load:
|
for x in to_load:
|
||||||
@ -62,6 +62,7 @@ def load_lora(lora, to_load):
|
|||||||
diffusers_lora = "{}_lora.up.weight".format(x)
|
diffusers_lora = "{}_lora.up.weight".format(x)
|
||||||
diffusers2_lora = "{}.lora_B.weight".format(x)
|
diffusers2_lora = "{}.lora_B.weight".format(x)
|
||||||
diffusers3_lora = "{}.lora.up.weight".format(x)
|
diffusers3_lora = "{}.lora.up.weight".format(x)
|
||||||
|
mochi_lora = "{}.lora_B".format(x)
|
||||||
transformers_lora = "{}.lora_linear_layer.up.weight".format(x)
|
transformers_lora = "{}.lora_linear_layer.up.weight".format(x)
|
||||||
A_name = None
|
A_name = None
|
||||||
|
|
||||||
@ -81,6 +82,10 @@ def load_lora(lora, to_load):
|
|||||||
A_name = diffusers3_lora
|
A_name = diffusers3_lora
|
||||||
B_name = "{}.lora.down.weight".format(x)
|
B_name = "{}.lora.down.weight".format(x)
|
||||||
mid_name = None
|
mid_name = None
|
||||||
|
elif mochi_lora in lora.keys():
|
||||||
|
A_name = mochi_lora
|
||||||
|
B_name = "{}.lora_A".format(x)
|
||||||
|
mid_name = None
|
||||||
elif transformers_lora in lora.keys():
|
elif transformers_lora in lora.keys():
|
||||||
A_name = transformers_lora
|
A_name = transformers_lora
|
||||||
B_name ="{}.lora_linear_layer.down.weight".format(x)
|
B_name ="{}.lora_linear_layer.down.weight".format(x)
|
||||||
@ -208,9 +213,10 @@ def load_lora(lora, to_load):
|
|||||||
patch_dict[to_load[x]] = ("set", (set_weight,))
|
patch_dict[to_load[x]] = ("set", (set_weight,))
|
||||||
loaded_keys.add(set_weight_name)
|
loaded_keys.add(set_weight_name)
|
||||||
|
|
||||||
for x in lora.keys():
|
if log_missing:
|
||||||
if x not in loaded_keys:
|
for x in lora.keys():
|
||||||
logging.warning("lora key not loaded: {}".format(x))
|
if x not in loaded_keys:
|
||||||
|
logging.warning("lora key not loaded: {}".format(x))
|
||||||
|
|
||||||
return patch_dict
|
return patch_dict
|
||||||
|
|
||||||
@ -362,6 +368,12 @@ def model_lora_keys_unet(model, key_map={}):
|
|||||||
key_map["lycoris_{}".format(k[:-len(".weight")].replace(".", "_"))] = to #simpletrainer lycoris
|
key_map["lycoris_{}".format(k[:-len(".weight")].replace(".", "_"))] = to #simpletrainer lycoris
|
||||||
key_map["lora_transformer_{}".format(k[:-len(".weight")].replace(".", "_"))] = to #onetrainer
|
key_map["lora_transformer_{}".format(k[:-len(".weight")].replace(".", "_"))] = to #onetrainer
|
||||||
|
|
||||||
|
if isinstance(model, comfy.model_base.GenmoMochi):
|
||||||
|
for k in sdk:
|
||||||
|
if k.startswith("diffusion_model.") and k.endswith(".weight"): #Official Mochi lora format
|
||||||
|
key_lora = k[len("diffusion_model."):-len(".weight")]
|
||||||
|
key_map["{}".format(key_lora)] = k
|
||||||
|
|
||||||
return key_map
|
return key_map
|
||||||
|
|
||||||
|
|
||||||
@ -418,7 +430,7 @@ def pad_tensor_to_shape(tensor: torch.Tensor, new_shape: list[int]) -> torch.Ten
|
|||||||
|
|
||||||
return padded_tensor
|
return padded_tensor
|
||||||
|
|
||||||
def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
|
def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32, original_weights=None):
|
||||||
for p in patches:
|
for p in patches:
|
||||||
strength = p[0]
|
strength = p[0]
|
||||||
v = p[1]
|
v = p[1]
|
||||||
@ -460,6 +472,11 @@ def calculate_weight(patches, weight, key, intermediate_dtype=torch.float32):
|
|||||||
weight += function(strength * comfy.model_management.cast_to_device(diff, weight.device, weight.dtype))
|
weight += function(strength * comfy.model_management.cast_to_device(diff, weight.device, weight.dtype))
|
||||||
elif patch_type == "set":
|
elif patch_type == "set":
|
||||||
weight.copy_(v[0])
|
weight.copy_(v[0])
|
||||||
|
elif patch_type == "model_as_lora":
|
||||||
|
target_weight: torch.Tensor = v[0]
|
||||||
|
diff_weight = comfy.model_management.cast_to_device(target_weight, weight.device, intermediate_dtype) - \
|
||||||
|
comfy.model_management.cast_to_device(original_weights[key][0][0], weight.device, intermediate_dtype)
|
||||||
|
weight += function(strength * comfy.model_management.cast_to_device(diff_weight, weight.device, weight.dtype))
|
||||||
elif patch_type == "lora": #lora/locon
|
elif patch_type == "lora": #lora/locon
|
||||||
mat1 = comfy.model_management.cast_to_device(v[0], weight.device, intermediate_dtype)
|
mat1 = comfy.model_management.cast_to_device(v[0], weight.device, intermediate_dtype)
|
||||||
mat2 = comfy.model_management.cast_to_device(v[1], weight.device, intermediate_dtype)
|
mat2 = comfy.model_management.cast_to_device(v[1], weight.device, intermediate_dtype)
|
||||||
|
|||||||
@ -33,12 +33,16 @@ import comfy.ldm.flux.model
|
|||||||
import comfy.ldm.lightricks.model
|
import comfy.ldm.lightricks.model
|
||||||
|
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
|
import comfy.patcher_extension
|
||||||
import comfy.conds
|
import comfy.conds
|
||||||
import comfy.ops
|
import comfy.ops
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
from . import utils
|
from . import utils
|
||||||
import comfy.latent_formats
|
import comfy.latent_formats
|
||||||
import math
|
import math
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from comfy.model_patcher import ModelPatcher
|
||||||
|
|
||||||
class ModelType(Enum):
|
class ModelType(Enum):
|
||||||
EPS = 1
|
EPS = 1
|
||||||
@ -95,6 +99,7 @@ class BaseModel(torch.nn.Module):
|
|||||||
self.model_config = model_config
|
self.model_config = model_config
|
||||||
self.manual_cast_dtype = model_config.manual_cast_dtype
|
self.manual_cast_dtype = model_config.manual_cast_dtype
|
||||||
self.device = device
|
self.device = device
|
||||||
|
self.current_patcher: 'ModelPatcher' = None
|
||||||
|
|
||||||
if not unet_config.get("disable_unet_model_creation", False):
|
if not unet_config.get("disable_unet_model_creation", False):
|
||||||
if model_config.custom_operations is None:
|
if model_config.custom_operations is None:
|
||||||
@ -120,6 +125,13 @@ class BaseModel(torch.nn.Module):
|
|||||||
self.memory_usage_factor = model_config.memory_usage_factor
|
self.memory_usage_factor = model_config.memory_usage_factor
|
||||||
|
|
||||||
def apply_model(self, x, t, c_concat=None, c_crossattn=None, control=None, transformer_options={}, **kwargs):
|
def apply_model(self, x, t, c_concat=None, c_crossattn=None, control=None, transformer_options={}, **kwargs):
|
||||||
|
return comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||||
|
self._apply_model,
|
||||||
|
self,
|
||||||
|
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.APPLY_MODEL, transformer_options)
|
||||||
|
).execute(x, t, c_concat, c_crossattn, control, transformer_options, **kwargs)
|
||||||
|
|
||||||
|
def _apply_model(self, x, t, c_concat=None, c_crossattn=None, control=None, transformer_options={}, **kwargs):
|
||||||
sigma = t
|
sigma = t
|
||||||
xc = self.model_sampling.calculate_input(sigma, x)
|
xc = self.model_sampling.calculate_input(sigma, x)
|
||||||
if c_concat is not None:
|
if c_concat is not None:
|
||||||
@ -712,7 +724,13 @@ class Flux(BaseModel):
|
|||||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.flux.model.Flux)
|
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.flux.model.Flux)
|
||||||
|
|
||||||
def concat_cond(self, **kwargs):
|
def concat_cond(self, **kwargs):
|
||||||
num_channels = self.diffusion_model.img_in.weight.shape[1] // (self.diffusion_model.patch_size * self.diffusion_model.patch_size)
|
try:
|
||||||
|
#Handle Flux control loras dynamically changing the img_in weight.
|
||||||
|
num_channels = self.diffusion_model.img_in.weight.shape[1] // (self.diffusion_model.patch_size * self.diffusion_model.patch_size)
|
||||||
|
except:
|
||||||
|
#Some cases like tensorrt might not have the weights accessible
|
||||||
|
num_channels = self.model_config.unet_config["in_channels"]
|
||||||
|
|
||||||
out_channels = self.model_config.unet_config["out_channels"]
|
out_channels = self.model_config.unet_config["out_channels"]
|
||||||
|
|
||||||
if num_channels <= out_channels:
|
if num_channels <= out_channels:
|
||||||
@ -786,5 +804,9 @@ class LTXV(BaseModel):
|
|||||||
if guiding_latent is not None:
|
if guiding_latent is not None:
|
||||||
out['guiding_latent'] = comfy.conds.CONDRegular(guiding_latent)
|
out['guiding_latent'] = comfy.conds.CONDRegular(guiding_latent)
|
||||||
|
|
||||||
|
guiding_latent_noise_scale = kwargs.get("guiding_latent_noise_scale", None)
|
||||||
|
if guiding_latent_noise_scale is not None:
|
||||||
|
out["guiding_latent_noise_scale"] = comfy.conds.CONDConstant(guiding_latent_noise_scale)
|
||||||
|
|
||||||
out['frame_rate'] = comfy.conds.CONDConstant(kwargs.get("frame_rate", 25))
|
out['frame_rate'] = comfy.conds.CONDConstant(kwargs.get("frame_rate", 25))
|
||||||
return out
|
return out
|
||||||
|
|||||||
@ -23,6 +23,8 @@ from comfy.cli_args import args
|
|||||||
import torch
|
import torch
|
||||||
import sys
|
import sys
|
||||||
import platform
|
import platform
|
||||||
|
import weakref
|
||||||
|
import gc
|
||||||
|
|
||||||
class VRAMState(Enum):
|
class VRAMState(Enum):
|
||||||
DISABLED = 0 #No vram present: no need to move models to vram
|
DISABLED = 0 #No vram present: no need to move models to vram
|
||||||
@ -287,11 +289,27 @@ def module_size(module):
|
|||||||
|
|
||||||
class LoadedModel:
|
class LoadedModel:
|
||||||
def __init__(self, model):
|
def __init__(self, model):
|
||||||
self.model = model
|
self._set_model(model)
|
||||||
self.device = model.load_device
|
self.device = model.load_device
|
||||||
self.weights_loaded = False
|
|
||||||
self.real_model = None
|
self.real_model = None
|
||||||
self.currently_used = True
|
self.currently_used = True
|
||||||
|
self.model_finalizer = None
|
||||||
|
self._patcher_finalizer = None
|
||||||
|
|
||||||
|
def _set_model(self, model):
|
||||||
|
self._model = weakref.ref(model)
|
||||||
|
if model.parent is not None:
|
||||||
|
self._parent_model = weakref.ref(model.parent)
|
||||||
|
self._patcher_finalizer = weakref.finalize(model, self._switch_parent)
|
||||||
|
|
||||||
|
def _switch_parent(self):
|
||||||
|
model = self._parent_model()
|
||||||
|
if model is not None:
|
||||||
|
self._set_model(model)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def model(self):
|
||||||
|
return self._model()
|
||||||
|
|
||||||
def model_memory(self):
|
def model_memory(self):
|
||||||
return self.model.model_size()
|
return self.model.model_size()
|
||||||
@ -306,32 +324,23 @@ class LoadedModel:
|
|||||||
return self.model_memory()
|
return self.model_memory()
|
||||||
|
|
||||||
def model_load(self, lowvram_model_memory=0, force_patch_weights=False):
|
def model_load(self, lowvram_model_memory=0, force_patch_weights=False):
|
||||||
patch_model_to = self.device
|
|
||||||
|
|
||||||
self.model.model_patches_to(self.device)
|
self.model.model_patches_to(self.device)
|
||||||
self.model.model_patches_to(self.model.model_dtype())
|
self.model.model_patches_to(self.model.model_dtype())
|
||||||
|
|
||||||
load_weights = not self.weights_loaded
|
# if self.model.loaded_size() > 0:
|
||||||
|
use_more_vram = lowvram_model_memory
|
||||||
|
if use_more_vram == 0:
|
||||||
|
use_more_vram = 1e32
|
||||||
|
self.model_use_more_vram(use_more_vram, force_patch_weights=force_patch_weights)
|
||||||
|
real_model = self.model.model
|
||||||
|
|
||||||
if self.model.loaded_size() > 0:
|
if is_intel_xpu() and not args.disable_ipex_optimize and 'ipex' in globals() and real_model is not None:
|
||||||
use_more_vram = lowvram_model_memory
|
|
||||||
if use_more_vram == 0:
|
|
||||||
use_more_vram = 1e32
|
|
||||||
self.model_use_more_vram(use_more_vram)
|
|
||||||
else:
|
|
||||||
try:
|
|
||||||
self.real_model = self.model.patch_model(device_to=patch_model_to, lowvram_model_memory=lowvram_model_memory, load_weights=load_weights, force_patch_weights=force_patch_weights)
|
|
||||||
except Exception as e:
|
|
||||||
self.model.unpatch_model(self.model.offload_device)
|
|
||||||
self.model_unload()
|
|
||||||
raise e
|
|
||||||
|
|
||||||
if is_intel_xpu() and not args.disable_ipex_optimize and 'ipex' in globals() and self.real_model is not None:
|
|
||||||
with torch.no_grad():
|
with torch.no_grad():
|
||||||
self.real_model = ipex.optimize(self.real_model.eval(), inplace=True, graph_mode=True, concat_linear=True)
|
real_model = ipex.optimize(real_model.eval(), inplace=True, graph_mode=True, concat_linear=True)
|
||||||
|
|
||||||
self.weights_loaded = True
|
self.real_model = weakref.ref(real_model)
|
||||||
return self.real_model
|
self.model_finalizer = weakref.finalize(real_model, cleanup_models)
|
||||||
|
return real_model
|
||||||
|
|
||||||
def should_reload_model(self, force_patch_weights=False):
|
def should_reload_model(self, force_patch_weights=False):
|
||||||
if force_patch_weights and self.model.lowvram_patch_counter() > 0:
|
if force_patch_weights and self.model.lowvram_patch_counter() > 0:
|
||||||
@ -344,18 +353,26 @@ class LoadedModel:
|
|||||||
freed = self.model.partially_unload(self.model.offload_device, memory_to_free)
|
freed = self.model.partially_unload(self.model.offload_device, memory_to_free)
|
||||||
if freed >= memory_to_free:
|
if freed >= memory_to_free:
|
||||||
return False
|
return False
|
||||||
self.model.unpatch_model(self.model.offload_device, unpatch_weights=unpatch_weights)
|
self.model.detach(unpatch_weights)
|
||||||
self.model.model_patches_to(self.model.offload_device)
|
self.model_finalizer.detach()
|
||||||
self.weights_loaded = self.weights_loaded and not unpatch_weights
|
self.model_finalizer = None
|
||||||
self.real_model = None
|
self.real_model = None
|
||||||
return True
|
return True
|
||||||
|
|
||||||
def model_use_more_vram(self, extra_memory):
|
def model_use_more_vram(self, extra_memory, force_patch_weights=False):
|
||||||
return self.model.partially_load(self.device, extra_memory)
|
return self.model.partially_load(self.device, extra_memory, force_patch_weights=force_patch_weights)
|
||||||
|
|
||||||
def __eq__(self, other):
|
def __eq__(self, other):
|
||||||
return self.model is other.model
|
return self.model is other.model
|
||||||
|
|
||||||
|
def __del__(self):
|
||||||
|
if self._patcher_finalizer is not None:
|
||||||
|
self._patcher_finalizer.detach()
|
||||||
|
|
||||||
|
def is_dead(self):
|
||||||
|
return self.real_model() is not None and self.model is None
|
||||||
|
|
||||||
|
|
||||||
def use_more_memory(extra_memory, loaded_models, device):
|
def use_more_memory(extra_memory, loaded_models, device):
|
||||||
for m in loaded_models:
|
for m in loaded_models:
|
||||||
if m.device == device:
|
if m.device == device:
|
||||||
@ -386,38 +403,8 @@ def extra_reserved_memory():
|
|||||||
def minimum_inference_memory():
|
def minimum_inference_memory():
|
||||||
return (1024 * 1024 * 1024) * 0.8 + extra_reserved_memory()
|
return (1024 * 1024 * 1024) * 0.8 + extra_reserved_memory()
|
||||||
|
|
||||||
def unload_model_clones(model, unload_weights_only=True, force_unload=True):
|
|
||||||
to_unload = []
|
|
||||||
for i in range(len(current_loaded_models)):
|
|
||||||
if model.is_clone(current_loaded_models[i].model):
|
|
||||||
to_unload = [i] + to_unload
|
|
||||||
|
|
||||||
if len(to_unload) == 0:
|
|
||||||
return True
|
|
||||||
|
|
||||||
same_weights = 0
|
|
||||||
for i in to_unload:
|
|
||||||
if model.clone_has_same_weights(current_loaded_models[i].model):
|
|
||||||
same_weights += 1
|
|
||||||
|
|
||||||
if same_weights == len(to_unload):
|
|
||||||
unload_weight = False
|
|
||||||
else:
|
|
||||||
unload_weight = True
|
|
||||||
|
|
||||||
if not force_unload:
|
|
||||||
if unload_weights_only and unload_weight == False:
|
|
||||||
return None
|
|
||||||
else:
|
|
||||||
unload_weight = True
|
|
||||||
|
|
||||||
for i in to_unload:
|
|
||||||
logging.debug("unload clone {} {}".format(i, unload_weight))
|
|
||||||
current_loaded_models.pop(i).model_unload(unpatch_weights=unload_weight)
|
|
||||||
|
|
||||||
return unload_weight
|
|
||||||
|
|
||||||
def free_memory(memory_required, device, keep_loaded=[]):
|
def free_memory(memory_required, device, keep_loaded=[]):
|
||||||
|
cleanup_models_gc()
|
||||||
unloaded_model = []
|
unloaded_model = []
|
||||||
can_unload = []
|
can_unload = []
|
||||||
unloaded_models = []
|
unloaded_models = []
|
||||||
@ -425,7 +412,7 @@ def free_memory(memory_required, device, keep_loaded=[]):
|
|||||||
for i in range(len(current_loaded_models) -1, -1, -1):
|
for i in range(len(current_loaded_models) -1, -1, -1):
|
||||||
shift_model = current_loaded_models[i]
|
shift_model = current_loaded_models[i]
|
||||||
if shift_model.device == device:
|
if shift_model.device == device:
|
||||||
if shift_model not in keep_loaded:
|
if shift_model not in keep_loaded and not shift_model.is_dead():
|
||||||
can_unload.append((-shift_model.model_offloaded_memory(), sys.getrefcount(shift_model.model), shift_model.model_memory(), i))
|
can_unload.append((-shift_model.model_offloaded_memory(), sys.getrefcount(shift_model.model), shift_model.model_memory(), i))
|
||||||
shift_model.currently_used = False
|
shift_model.currently_used = False
|
||||||
|
|
||||||
@ -454,6 +441,7 @@ def free_memory(memory_required, device, keep_loaded=[]):
|
|||||||
return unloaded_models
|
return unloaded_models
|
||||||
|
|
||||||
def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimum_memory_required=None, force_full_load=False):
|
def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimum_memory_required=None, force_full_load=False):
|
||||||
|
cleanup_models_gc()
|
||||||
global vram_state
|
global vram_state
|
||||||
|
|
||||||
inference_memory = minimum_inference_memory()
|
inference_memory = minimum_inference_memory()
|
||||||
@ -466,11 +454,9 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu
|
|||||||
models = set(models)
|
models = set(models)
|
||||||
|
|
||||||
models_to_load = []
|
models_to_load = []
|
||||||
models_already_loaded = []
|
|
||||||
for x in models:
|
for x in models:
|
||||||
loaded_model = LoadedModel(x)
|
loaded_model = LoadedModel(x)
|
||||||
loaded = None
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
loaded_model_index = current_loaded_models.index(loaded_model)
|
loaded_model_index = current_loaded_models.index(loaded_model)
|
||||||
except:
|
except:
|
||||||
@ -478,51 +464,35 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu
|
|||||||
|
|
||||||
if loaded_model_index is not None:
|
if loaded_model_index is not None:
|
||||||
loaded = current_loaded_models[loaded_model_index]
|
loaded = current_loaded_models[loaded_model_index]
|
||||||
if loaded.should_reload_model(force_patch_weights=force_patch_weights): #TODO: cleanup this model reload logic
|
loaded.currently_used = True
|
||||||
current_loaded_models.pop(loaded_model_index).model_unload(unpatch_weights=True)
|
models_to_load.append(loaded)
|
||||||
loaded = None
|
else:
|
||||||
else:
|
|
||||||
loaded.currently_used = True
|
|
||||||
models_already_loaded.append(loaded)
|
|
||||||
|
|
||||||
if loaded is None:
|
|
||||||
if hasattr(x, "model"):
|
if hasattr(x, "model"):
|
||||||
logging.info(f"Requested to load {x.model.__class__.__name__}")
|
logging.info(f"Requested to load {x.model.__class__.__name__}")
|
||||||
models_to_load.append(loaded_model)
|
models_to_load.append(loaded_model)
|
||||||
|
|
||||||
if len(models_to_load) == 0:
|
for loaded_model in models_to_load:
|
||||||
devs = set(map(lambda a: a.device, models_already_loaded))
|
to_unload = []
|
||||||
for d in devs:
|
for i in range(len(current_loaded_models)):
|
||||||
if d != torch.device("cpu"):
|
if loaded_model.model.is_clone(current_loaded_models[i].model):
|
||||||
free_memory(extra_mem + offloaded_memory(models_already_loaded, d), d, models_already_loaded)
|
to_unload = [i] + to_unload
|
||||||
free_mem = get_free_memory(d)
|
for i in to_unload:
|
||||||
if free_mem < minimum_memory_required:
|
current_loaded_models.pop(i).model.detach(unpatch_all=False)
|
||||||
logging.info("Unloading models for lowram load.") #TODO: partial model unloading when this case happens, also handle the opposite case where models can be unlowvramed.
|
|
||||||
models_to_load = free_memory(minimum_memory_required, d)
|
|
||||||
logging.info("{} models unloaded.".format(len(models_to_load)))
|
|
||||||
else:
|
|
||||||
use_more_memory(free_mem - minimum_memory_required, models_already_loaded, d)
|
|
||||||
if len(models_to_load) == 0:
|
|
||||||
return
|
|
||||||
|
|
||||||
logging.info(f"Loading {len(models_to_load)} new model{'s' if len(models_to_load) > 1 else ''}")
|
|
||||||
|
|
||||||
total_memory_required = {}
|
total_memory_required = {}
|
||||||
for loaded_model in models_to_load:
|
for loaded_model in models_to_load:
|
||||||
unload_model_clones(loaded_model.model, unload_weights_only=True, force_unload=False) #unload clones where the weights are different
|
|
||||||
total_memory_required[loaded_model.device] = total_memory_required.get(loaded_model.device, 0) + loaded_model.model_memory_required(loaded_model.device)
|
total_memory_required[loaded_model.device] = total_memory_required.get(loaded_model.device, 0) + loaded_model.model_memory_required(loaded_model.device)
|
||||||
|
|
||||||
for loaded_model in models_already_loaded:
|
|
||||||
total_memory_required[loaded_model.device] = total_memory_required.get(loaded_model.device, 0) + loaded_model.model_memory_required(loaded_model.device)
|
|
||||||
|
|
||||||
for loaded_model in models_to_load:
|
|
||||||
weights_unloaded = unload_model_clones(loaded_model.model, unload_weights_only=False, force_unload=False) #unload the rest of the clones where the weights can stay loaded
|
|
||||||
if weights_unloaded is not None:
|
|
||||||
loaded_model.weights_loaded = not weights_unloaded
|
|
||||||
|
|
||||||
for device in total_memory_required:
|
for device in total_memory_required:
|
||||||
if device != torch.device("cpu"):
|
if device != torch.device("cpu"):
|
||||||
free_memory(total_memory_required[device] * 1.1 + extra_mem, device, models_already_loaded)
|
free_memory(total_memory_required[device] * 1.1 + extra_mem, device)
|
||||||
|
|
||||||
|
for device in total_memory_required:
|
||||||
|
if device != torch.device("cpu"):
|
||||||
|
free_mem = get_free_memory(device)
|
||||||
|
if free_mem < minimum_memory_required:
|
||||||
|
models_l = free_memory(minimum_memory_required, device)
|
||||||
|
logging.info("{} models unloaded.".format(len(models_l)))
|
||||||
|
|
||||||
for loaded_model in models_to_load:
|
for loaded_model in models_to_load:
|
||||||
model = loaded_model.model
|
model = loaded_model.model
|
||||||
@ -544,17 +514,8 @@ def load_models_gpu(models, memory_required=0, force_patch_weights=False, minimu
|
|||||||
|
|
||||||
cur_loaded_model = loaded_model.model_load(lowvram_model_memory, force_patch_weights=force_patch_weights)
|
cur_loaded_model = loaded_model.model_load(lowvram_model_memory, force_patch_weights=force_patch_weights)
|
||||||
current_loaded_models.insert(0, loaded_model)
|
current_loaded_models.insert(0, loaded_model)
|
||||||
|
|
||||||
|
|
||||||
devs = set(map(lambda a: a.device, models_already_loaded))
|
|
||||||
for d in devs:
|
|
||||||
if d != torch.device("cpu"):
|
|
||||||
free_mem = get_free_memory(d)
|
|
||||||
if free_mem > minimum_memory_required:
|
|
||||||
use_more_memory(free_mem - minimum_memory_required, models_already_loaded, d)
|
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|
||||||
def load_model_gpu(model):
|
def load_model_gpu(model):
|
||||||
return load_models_gpu([model])
|
return load_models_gpu([model])
|
||||||
|
|
||||||
@ -568,21 +529,35 @@ def loaded_models(only_currently_used=False):
|
|||||||
output.append(m.model)
|
output.append(m.model)
|
||||||
return output
|
return output
|
||||||
|
|
||||||
def cleanup_models(keep_clone_weights_loaded=False):
|
|
||||||
|
def cleanup_models_gc():
|
||||||
|
do_gc = False
|
||||||
|
for i in range(len(current_loaded_models)):
|
||||||
|
cur = current_loaded_models[i]
|
||||||
|
if cur.is_dead():
|
||||||
|
logging.info("Potential memory leak detected with model {}, doing a full garbage collect, for maximum performance avoid circular references in the model code.".format(cur.real_model().__class__.__name__))
|
||||||
|
do_gc = True
|
||||||
|
break
|
||||||
|
|
||||||
|
if do_gc:
|
||||||
|
gc.collect()
|
||||||
|
soft_empty_cache()
|
||||||
|
|
||||||
|
for i in range(len(current_loaded_models)):
|
||||||
|
cur = current_loaded_models[i]
|
||||||
|
if cur.is_dead():
|
||||||
|
logging.warning("WARNING, memory leak with model {}. Please make sure it is not being referenced from somewhere.".format(cur.real_model().__class__.__name__))
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
def cleanup_models():
|
||||||
to_delete = []
|
to_delete = []
|
||||||
for i in range(len(current_loaded_models)):
|
for i in range(len(current_loaded_models)):
|
||||||
#TODO: very fragile function needs improvement
|
if current_loaded_models[i].real_model() is None:
|
||||||
num_refs = sys.getrefcount(current_loaded_models[i].model)
|
to_delete = [i] + to_delete
|
||||||
if num_refs <= 2:
|
|
||||||
if not keep_clone_weights_loaded:
|
|
||||||
to_delete = [i] + to_delete
|
|
||||||
#TODO: find a less fragile way to do this.
|
|
||||||
elif sys.getrefcount(current_loaded_models[i].real_model) <= 3: #references from .real_model + the .model
|
|
||||||
to_delete = [i] + to_delete
|
|
||||||
|
|
||||||
for i in to_delete:
|
for i in to_delete:
|
||||||
x = current_loaded_models.pop(i)
|
x = current_loaded_models.pop(i)
|
||||||
x.model_unload()
|
|
||||||
del x
|
del x
|
||||||
|
|
||||||
def dtype_size(dtype):
|
def dtype_size(dtype):
|
||||||
@ -628,6 +603,10 @@ def maximum_vram_for_weights(device=None):
|
|||||||
def unet_dtype(device=None, model_params=0, supported_dtypes=[torch.float16, torch.bfloat16, torch.float32]):
|
def unet_dtype(device=None, model_params=0, supported_dtypes=[torch.float16, torch.bfloat16, torch.float32]):
|
||||||
if model_params < 0:
|
if model_params < 0:
|
||||||
model_params = 1000000000000000000000
|
model_params = 1000000000000000000000
|
||||||
|
if args.fp32_unet:
|
||||||
|
return torch.float32
|
||||||
|
if args.fp64_unet:
|
||||||
|
return torch.float64
|
||||||
if args.bf16_unet:
|
if args.bf16_unet:
|
||||||
return torch.bfloat16
|
return torch.bfloat16
|
||||||
if args.fp16_unet:
|
if args.fp16_unet:
|
||||||
@ -674,7 +653,7 @@ def unet_dtype(device=None, model_params=0, supported_dtypes=[torch.float16, tor
|
|||||||
|
|
||||||
# None means no manual cast
|
# None means no manual cast
|
||||||
def unet_manual_cast(weight_dtype, inference_device, supported_dtypes=[torch.float16, torch.bfloat16, torch.float32]):
|
def unet_manual_cast(weight_dtype, inference_device, supported_dtypes=[torch.float16, torch.bfloat16, torch.float32]):
|
||||||
if weight_dtype == torch.float32:
|
if weight_dtype == torch.float32 or weight_dtype == torch.float64:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
fp16_supported = should_use_fp16(inference_device, prioritize_performance=False)
|
fp16_supported = should_use_fp16(inference_device, prioritize_performance=False)
|
||||||
|
|||||||
@ -16,6 +16,8 @@
|
|||||||
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
along with this program. If not, see <https://www.gnu.org/licenses/>.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
from typing import Optional, Callable
|
||||||
import torch
|
import torch
|
||||||
import copy
|
import copy
|
||||||
import inspect
|
import inspect
|
||||||
@ -28,6 +30,9 @@ import comfy.utils
|
|||||||
import comfy.float
|
import comfy.float
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
import comfy.lora
|
import comfy.lora
|
||||||
|
import comfy.hooks
|
||||||
|
import comfy.patcher_extension
|
||||||
|
from comfy.patcher_extension import CallbacksMP, WrappersMP, PatcherInjection
|
||||||
from comfy.comfy_types import UnetWrapperFunction
|
from comfy.comfy_types import UnetWrapperFunction
|
||||||
|
|
||||||
def string_to_seed(data):
|
def string_to_seed(data):
|
||||||
@ -76,6 +81,17 @@ def set_model_options_pre_cfg_function(model_options, pre_cfg_function, disable_
|
|||||||
model_options["disable_cfg1_optimization"] = True
|
model_options["disable_cfg1_optimization"] = True
|
||||||
return model_options
|
return model_options
|
||||||
|
|
||||||
|
def create_model_options_clone(orig_model_options: dict):
|
||||||
|
return comfy.patcher_extension.copy_nested_dicts(orig_model_options)
|
||||||
|
|
||||||
|
def create_hook_patches_clone(orig_hook_patches):
|
||||||
|
new_hook_patches = {}
|
||||||
|
for hook_ref in orig_hook_patches:
|
||||||
|
new_hook_patches[hook_ref] = {}
|
||||||
|
for k in orig_hook_patches[hook_ref]:
|
||||||
|
new_hook_patches[hook_ref][k] = orig_hook_patches[hook_ref][k][:]
|
||||||
|
return new_hook_patches
|
||||||
|
|
||||||
def wipe_lowvram_weight(m):
|
def wipe_lowvram_weight(m):
|
||||||
if hasattr(m, "prev_comfy_cast_weights"):
|
if hasattr(m, "prev_comfy_cast_weights"):
|
||||||
m.comfy_cast_weights = m.prev_comfy_cast_weights
|
m.comfy_cast_weights = m.prev_comfy_cast_weights
|
||||||
@ -119,6 +135,49 @@ def get_key_weight(model, key):
|
|||||||
|
|
||||||
return weight, set_func, convert_func
|
return weight, set_func, convert_func
|
||||||
|
|
||||||
|
class AutoPatcherEjector:
|
||||||
|
def __init__(self, model: 'ModelPatcher', skip_and_inject_on_exit_only=False):
|
||||||
|
self.model = model
|
||||||
|
self.was_injected = False
|
||||||
|
self.prev_skip_injection = False
|
||||||
|
self.skip_and_inject_on_exit_only = skip_and_inject_on_exit_only
|
||||||
|
|
||||||
|
def __enter__(self):
|
||||||
|
self.was_injected = False
|
||||||
|
self.prev_skip_injection = self.model.skip_injection
|
||||||
|
if self.skip_and_inject_on_exit_only:
|
||||||
|
self.model.skip_injection = True
|
||||||
|
if self.model.is_injected:
|
||||||
|
self.model.eject_model()
|
||||||
|
self.was_injected = True
|
||||||
|
|
||||||
|
def __exit__(self, *args):
|
||||||
|
if self.skip_and_inject_on_exit_only:
|
||||||
|
self.model.skip_injection = self.prev_skip_injection
|
||||||
|
self.model.inject_model()
|
||||||
|
if self.was_injected and not self.model.skip_injection:
|
||||||
|
self.model.inject_model()
|
||||||
|
self.model.skip_injection = self.prev_skip_injection
|
||||||
|
|
||||||
|
class MemoryCounter:
|
||||||
|
def __init__(self, initial: int, minimum=0):
|
||||||
|
self.value = initial
|
||||||
|
self.minimum = minimum
|
||||||
|
# TODO: add a safe limit besides 0
|
||||||
|
|
||||||
|
def use(self, weight: torch.Tensor):
|
||||||
|
weight_size = weight.nelement() * weight.element_size()
|
||||||
|
if self.is_useable(weight_size):
|
||||||
|
self.decrement(weight_size)
|
||||||
|
return True
|
||||||
|
return False
|
||||||
|
|
||||||
|
def is_useable(self, used: int):
|
||||||
|
return self.value - used > self.minimum
|
||||||
|
|
||||||
|
def decrement(self, used: int):
|
||||||
|
self.value -= used
|
||||||
|
|
||||||
class ModelPatcher:
|
class ModelPatcher:
|
||||||
def __init__(self, model, load_device, offload_device, size=0, weight_inplace_update=False):
|
def __init__(self, model, load_device, offload_device, size=0, weight_inplace_update=False):
|
||||||
self.size = size
|
self.size = size
|
||||||
@ -139,6 +198,25 @@ class ModelPatcher:
|
|||||||
self.offload_device = offload_device
|
self.offload_device = offload_device
|
||||||
self.weight_inplace_update = weight_inplace_update
|
self.weight_inplace_update = weight_inplace_update
|
||||||
self.patches_uuid = uuid.uuid4()
|
self.patches_uuid = uuid.uuid4()
|
||||||
|
self.parent = None
|
||||||
|
|
||||||
|
self.attachments: dict[str] = {}
|
||||||
|
self.additional_models: dict[str, list[ModelPatcher]] = {}
|
||||||
|
self.callbacks: dict[str, dict[str, list[Callable]]] = CallbacksMP.init_callbacks()
|
||||||
|
self.wrappers: dict[str, dict[str, list[Callable]]] = WrappersMP.init_wrappers()
|
||||||
|
|
||||||
|
self.is_injected = False
|
||||||
|
self.skip_injection = False
|
||||||
|
self.injections: dict[str, list[PatcherInjection]] = {}
|
||||||
|
|
||||||
|
self.hook_patches: dict[comfy.hooks._HookRef] = {}
|
||||||
|
self.hook_patches_backup: dict[comfy.hooks._HookRef] = {}
|
||||||
|
self.hook_backup: dict[str, tuple[torch.Tensor, torch.device]] = {}
|
||||||
|
self.cached_hook_patches: dict[comfy.hooks.HookGroup, dict[str, torch.Tensor]] = {}
|
||||||
|
self.current_hooks: Optional[comfy.hooks.HookGroup] = None
|
||||||
|
self.forced_hooks: Optional[comfy.hooks.HookGroup] = None # NOTE: only used for CLIP at this time
|
||||||
|
self.is_clip = False
|
||||||
|
self.hook_mode = comfy.hooks.EnumHookMode.MaxSpeed
|
||||||
|
|
||||||
if not hasattr(self.model, 'model_loaded_weight_memory'):
|
if not hasattr(self.model, 'model_loaded_weight_memory'):
|
||||||
self.model.model_loaded_weight_memory = 0
|
self.model.model_loaded_weight_memory = 0
|
||||||
@ -149,6 +227,9 @@ class ModelPatcher:
|
|||||||
if not hasattr(self.model, 'model_lowvram'):
|
if not hasattr(self.model, 'model_lowvram'):
|
||||||
self.model.model_lowvram = False
|
self.model.model_lowvram = False
|
||||||
|
|
||||||
|
if not hasattr(self.model, 'current_weight_patches_uuid'):
|
||||||
|
self.model.current_weight_patches_uuid = None
|
||||||
|
|
||||||
def model_size(self):
|
def model_size(self):
|
||||||
if self.size > 0:
|
if self.size > 0:
|
||||||
return self.size
|
return self.size
|
||||||
@ -162,7 +243,7 @@ class ModelPatcher:
|
|||||||
return self.model.lowvram_patch_counter
|
return self.model.lowvram_patch_counter
|
||||||
|
|
||||||
def clone(self):
|
def clone(self):
|
||||||
n = ModelPatcher(self.model, self.load_device, self.offload_device, self.size, weight_inplace_update=self.weight_inplace_update)
|
n = self.__class__(self.model, self.load_device, self.offload_device, self.size, weight_inplace_update=self.weight_inplace_update)
|
||||||
n.patches = {}
|
n.patches = {}
|
||||||
for k in self.patches:
|
for k in self.patches:
|
||||||
n.patches[k] = self.patches[k][:]
|
n.patches[k] = self.patches[k][:]
|
||||||
@ -172,6 +253,48 @@ class ModelPatcher:
|
|||||||
n.model_options = copy.deepcopy(self.model_options)
|
n.model_options = copy.deepcopy(self.model_options)
|
||||||
n.backup = self.backup
|
n.backup = self.backup
|
||||||
n.object_patches_backup = self.object_patches_backup
|
n.object_patches_backup = self.object_patches_backup
|
||||||
|
n.parent = self
|
||||||
|
|
||||||
|
# attachments
|
||||||
|
n.attachments = {}
|
||||||
|
for k in self.attachments:
|
||||||
|
if hasattr(self.attachments[k], "on_model_patcher_clone"):
|
||||||
|
n.attachments[k] = self.attachments[k].on_model_patcher_clone()
|
||||||
|
else:
|
||||||
|
n.attachments[k] = self.attachments[k]
|
||||||
|
# additional models
|
||||||
|
for k, c in self.additional_models.items():
|
||||||
|
n.additional_models[k] = [x.clone() for x in c]
|
||||||
|
# callbacks
|
||||||
|
for k, c in self.callbacks.items():
|
||||||
|
n.callbacks[k] = {}
|
||||||
|
for k1, c1 in c.items():
|
||||||
|
n.callbacks[k][k1] = c1.copy()
|
||||||
|
# sample wrappers
|
||||||
|
for k, w in self.wrappers.items():
|
||||||
|
n.wrappers[k] = {}
|
||||||
|
for k1, w1 in w.items():
|
||||||
|
n.wrappers[k][k1] = w1.copy()
|
||||||
|
# injection
|
||||||
|
n.is_injected = self.is_injected
|
||||||
|
n.skip_injection = self.skip_injection
|
||||||
|
for k, i in self.injections.items():
|
||||||
|
n.injections[k] = i.copy()
|
||||||
|
# hooks
|
||||||
|
n.hook_patches = create_hook_patches_clone(self.hook_patches)
|
||||||
|
n.hook_patches_backup = create_hook_patches_clone(self.hook_patches_backup)
|
||||||
|
for group in self.cached_hook_patches:
|
||||||
|
n.cached_hook_patches[group] = {}
|
||||||
|
for k in self.cached_hook_patches[group]:
|
||||||
|
n.cached_hook_patches[group][k] = self.cached_hook_patches[group][k]
|
||||||
|
n.hook_backup = self.hook_backup
|
||||||
|
n.current_hooks = self.current_hooks.clone() if self.current_hooks else self.current_hooks
|
||||||
|
n.forced_hooks = self.forced_hooks.clone() if self.forced_hooks else self.forced_hooks
|
||||||
|
n.is_clip = self.is_clip
|
||||||
|
n.hook_mode = self.hook_mode
|
||||||
|
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_CLONE):
|
||||||
|
callback(self, n)
|
||||||
return n
|
return n
|
||||||
|
|
||||||
def is_clone(self, other):
|
def is_clone(self, other):
|
||||||
@ -179,10 +302,29 @@ class ModelPatcher:
|
|||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
def clone_has_same_weights(self, clone):
|
def clone_has_same_weights(self, clone: 'ModelPatcher'):
|
||||||
if not self.is_clone(clone):
|
if not self.is_clone(clone):
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
if self.current_hooks != clone.current_hooks:
|
||||||
|
return False
|
||||||
|
if self.forced_hooks != clone.forced_hooks:
|
||||||
|
return False
|
||||||
|
if self.hook_patches.keys() != clone.hook_patches.keys():
|
||||||
|
return False
|
||||||
|
if self.attachments.keys() != clone.attachments.keys():
|
||||||
|
return False
|
||||||
|
if self.additional_models.keys() != clone.additional_models.keys():
|
||||||
|
return False
|
||||||
|
for key in self.callbacks:
|
||||||
|
if len(self.callbacks[key]) != len(clone.callbacks[key]):
|
||||||
|
return False
|
||||||
|
for key in self.wrappers:
|
||||||
|
if len(self.wrappers[key]) != len(clone.wrappers[key]):
|
||||||
|
return False
|
||||||
|
if self.injections.keys() != clone.injections.keys():
|
||||||
|
return False
|
||||||
|
|
||||||
if len(self.patches) == 0 and len(clone.patches) == 0:
|
if len(self.patches) == 0 and len(clone.patches) == 0:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
@ -251,6 +393,12 @@ class ModelPatcher:
|
|||||||
def set_model_output_block_patch(self, patch):
|
def set_model_output_block_patch(self, patch):
|
||||||
self.set_model_patch(patch, "output_block_patch")
|
self.set_model_patch(patch, "output_block_patch")
|
||||||
|
|
||||||
|
def set_model_emb_patch(self, patch):
|
||||||
|
self.set_model_patch(patch, "emb_patch")
|
||||||
|
|
||||||
|
def set_model_forward_timestep_embed_patch(self, patch):
|
||||||
|
self.set_model_patch(patch, "forward_timestep_embed_patch")
|
||||||
|
|
||||||
def add_object_patch(self, name, obj):
|
def add_object_patch(self, name, obj):
|
||||||
self.object_patches[name] = obj
|
self.object_patches[name] = obj
|
||||||
|
|
||||||
@ -289,27 +437,28 @@ class ModelPatcher:
|
|||||||
return self.model.get_dtype()
|
return self.model.get_dtype()
|
||||||
|
|
||||||
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
||||||
p = set()
|
with self.use_ejected():
|
||||||
model_sd = self.model.state_dict()
|
p = set()
|
||||||
for k in patches:
|
model_sd = self.model.state_dict()
|
||||||
offset = None
|
for k in patches:
|
||||||
function = None
|
offset = None
|
||||||
if isinstance(k, str):
|
function = None
|
||||||
key = k
|
if isinstance(k, str):
|
||||||
else:
|
key = k
|
||||||
offset = k[1]
|
else:
|
||||||
key = k[0]
|
offset = k[1]
|
||||||
if len(k) > 2:
|
key = k[0]
|
||||||
function = k[2]
|
if len(k) > 2:
|
||||||
|
function = k[2]
|
||||||
|
|
||||||
if key in model_sd:
|
if key in model_sd:
|
||||||
p.add(k)
|
p.add(k)
|
||||||
current_patches = self.patches.get(key, [])
|
current_patches = self.patches.get(key, [])
|
||||||
current_patches.append((strength_patch, patches[k], strength_model, offset, function))
|
current_patches.append((strength_patch, patches[k], strength_model, offset, function))
|
||||||
self.patches[key] = current_patches
|
self.patches[key] = current_patches
|
||||||
|
|
||||||
self.patches_uuid = uuid.uuid4()
|
self.patches_uuid = uuid.uuid4()
|
||||||
return list(p)
|
return list(p)
|
||||||
|
|
||||||
def get_key_patches(self, filter_prefix=None):
|
def get_key_patches(self, filter_prefix=None):
|
||||||
model_sd = self.model_state_dict()
|
model_sd = self.model_state_dict()
|
||||||
@ -319,9 +468,12 @@ class ModelPatcher:
|
|||||||
if not k.startswith(filter_prefix):
|
if not k.startswith(filter_prefix):
|
||||||
continue
|
continue
|
||||||
bk = self.backup.get(k, None)
|
bk = self.backup.get(k, None)
|
||||||
|
hbk = self.hook_backup.get(k, None)
|
||||||
weight, set_func, convert_func = get_key_weight(self.model, k)
|
weight, set_func, convert_func = get_key_weight(self.model, k)
|
||||||
if bk is not None:
|
if bk is not None:
|
||||||
weight = bk.weight
|
weight = bk.weight
|
||||||
|
if hbk is not None:
|
||||||
|
weight = hbk[0]
|
||||||
if convert_func is None:
|
if convert_func is None:
|
||||||
convert_func = lambda a, **kwargs: a
|
convert_func = lambda a, **kwargs: a
|
||||||
|
|
||||||
@ -332,13 +484,14 @@ class ModelPatcher:
|
|||||||
return p
|
return p
|
||||||
|
|
||||||
def model_state_dict(self, filter_prefix=None):
|
def model_state_dict(self, filter_prefix=None):
|
||||||
sd = self.model.state_dict()
|
with self.use_ejected():
|
||||||
keys = list(sd.keys())
|
sd = self.model.state_dict()
|
||||||
if filter_prefix is not None:
|
keys = list(sd.keys())
|
||||||
for k in keys:
|
if filter_prefix is not None:
|
||||||
if not k.startswith(filter_prefix):
|
for k in keys:
|
||||||
sd.pop(k)
|
if not k.startswith(filter_prefix):
|
||||||
return sd
|
sd.pop(k)
|
||||||
|
return sd
|
||||||
|
|
||||||
def patch_weight_to_device(self, key, device_to=None, inplace_update=False):
|
def patch_weight_to_device(self, key, device_to=None, inplace_update=False):
|
||||||
if key not in self.patches:
|
if key not in self.patches:
|
||||||
@ -383,105 +536,117 @@ class ModelPatcher:
|
|||||||
return loading
|
return loading
|
||||||
|
|
||||||
def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False):
|
def load(self, device_to=None, lowvram_model_memory=0, force_patch_weights=False, full_load=False):
|
||||||
mem_counter = 0
|
with self.use_ejected():
|
||||||
patch_counter = 0
|
self.unpatch_hooks()
|
||||||
lowvram_counter = 0
|
mem_counter = 0
|
||||||
loading = self._load_list()
|
patch_counter = 0
|
||||||
|
lowvram_counter = 0
|
||||||
|
loading = self._load_list()
|
||||||
|
|
||||||
load_completely = []
|
load_completely = []
|
||||||
loading.sort(reverse=True)
|
loading.sort(reverse=True)
|
||||||
for x in loading:
|
for x in loading:
|
||||||
n = x[1]
|
n = x[1]
|
||||||
m = x[2]
|
m = x[2]
|
||||||
params = x[3]
|
params = x[3]
|
||||||
module_mem = x[0]
|
module_mem = x[0]
|
||||||
|
|
||||||
lowvram_weight = False
|
lowvram_weight = False
|
||||||
|
|
||||||
if not full_load and hasattr(m, "comfy_cast_weights"):
|
if not full_load and hasattr(m, "comfy_cast_weights"):
|
||||||
if mem_counter + module_mem >= lowvram_model_memory:
|
if mem_counter + module_mem >= lowvram_model_memory:
|
||||||
lowvram_weight = True
|
lowvram_weight = True
|
||||||
lowvram_counter += 1
|
lowvram_counter += 1
|
||||||
if hasattr(m, "prev_comfy_cast_weights"): #Already lowvramed
|
if hasattr(m, "prev_comfy_cast_weights"): #Already lowvramed
|
||||||
|
continue
|
||||||
|
|
||||||
|
weight_key = "{}.weight".format(n)
|
||||||
|
bias_key = "{}.bias".format(n)
|
||||||
|
|
||||||
|
if lowvram_weight:
|
||||||
|
if weight_key in self.patches:
|
||||||
|
if force_patch_weights:
|
||||||
|
self.patch_weight_to_device(weight_key)
|
||||||
|
else:
|
||||||
|
m.weight_function = LowVramPatch(weight_key, self.patches)
|
||||||
|
patch_counter += 1
|
||||||
|
if bias_key in self.patches:
|
||||||
|
if force_patch_weights:
|
||||||
|
self.patch_weight_to_device(bias_key)
|
||||||
|
else:
|
||||||
|
m.bias_function = LowVramPatch(bias_key, self.patches)
|
||||||
|
patch_counter += 1
|
||||||
|
|
||||||
|
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
||||||
|
m.comfy_cast_weights = True
|
||||||
|
else:
|
||||||
|
if hasattr(m, "comfy_cast_weights"):
|
||||||
|
if m.comfy_cast_weights:
|
||||||
|
wipe_lowvram_weight(m)
|
||||||
|
|
||||||
|
if full_load or mem_counter + module_mem < lowvram_model_memory:
|
||||||
|
mem_counter += module_mem
|
||||||
|
load_completely.append((module_mem, n, m, params))
|
||||||
|
|
||||||
|
load_completely.sort(reverse=True)
|
||||||
|
for x in load_completely:
|
||||||
|
n = x[1]
|
||||||
|
m = x[2]
|
||||||
|
params = x[3]
|
||||||
|
if hasattr(m, "comfy_patched_weights"):
|
||||||
|
if m.comfy_patched_weights == True:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
weight_key = "{}.weight".format(n)
|
for param in params:
|
||||||
bias_key = "{}.bias".format(n)
|
self.patch_weight_to_device("{}.{}".format(n, param), device_to=device_to)
|
||||||
|
|
||||||
if lowvram_weight:
|
logging.debug("lowvram: loaded module regularly {} {}".format(n, m))
|
||||||
if weight_key in self.patches:
|
m.comfy_patched_weights = True
|
||||||
if force_patch_weights:
|
|
||||||
self.patch_weight_to_device(weight_key)
|
|
||||||
else:
|
|
||||||
m.weight_function = LowVramPatch(weight_key, self.patches)
|
|
||||||
patch_counter += 1
|
|
||||||
if bias_key in self.patches:
|
|
||||||
if force_patch_weights:
|
|
||||||
self.patch_weight_to_device(bias_key)
|
|
||||||
else:
|
|
||||||
m.bias_function = LowVramPatch(bias_key, self.patches)
|
|
||||||
patch_counter += 1
|
|
||||||
|
|
||||||
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
for x in load_completely:
|
||||||
m.comfy_cast_weights = True
|
x[2].to(device_to)
|
||||||
|
|
||||||
|
if lowvram_counter > 0:
|
||||||
|
logging.info("loaded partially {} {} {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), patch_counter))
|
||||||
|
self.model.model_lowvram = True
|
||||||
else:
|
else:
|
||||||
if hasattr(m, "comfy_cast_weights"):
|
logging.info("loaded completely {} {} {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), full_load))
|
||||||
if m.comfy_cast_weights:
|
self.model.model_lowvram = False
|
||||||
wipe_lowvram_weight(m)
|
if full_load:
|
||||||
|
self.model.to(device_to)
|
||||||
|
mem_counter = self.model_size()
|
||||||
|
|
||||||
if full_load or mem_counter + module_mem < lowvram_model_memory:
|
self.model.lowvram_patch_counter += patch_counter
|
||||||
mem_counter += module_mem
|
self.model.device = device_to
|
||||||
load_completely.append((module_mem, n, m, params))
|
self.model.model_loaded_weight_memory = mem_counter
|
||||||
|
self.model.current_weight_patches_uuid = self.patches_uuid
|
||||||
|
|
||||||
load_completely.sort(reverse=True)
|
for callback in self.get_all_callbacks(CallbacksMP.ON_LOAD):
|
||||||
for x in load_completely:
|
callback(self, device_to, lowvram_model_memory, force_patch_weights, full_load)
|
||||||
n = x[1]
|
|
||||||
m = x[2]
|
|
||||||
params = x[3]
|
|
||||||
if hasattr(m, "comfy_patched_weights"):
|
|
||||||
if m.comfy_patched_weights == True:
|
|
||||||
continue
|
|
||||||
|
|
||||||
for param in params:
|
self.apply_hooks(self.forced_hooks, force_apply=True)
|
||||||
self.patch_weight_to_device("{}.{}".format(n, param), device_to=device_to)
|
|
||||||
|
|
||||||
logging.debug("lowvram: loaded module regularly {} {}".format(n, m))
|
|
||||||
m.comfy_patched_weights = True
|
|
||||||
|
|
||||||
for x in load_completely:
|
|
||||||
x[2].to(device_to)
|
|
||||||
|
|
||||||
if lowvram_counter > 0:
|
|
||||||
logging.info("loaded partially {} {} {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), patch_counter))
|
|
||||||
self.model.model_lowvram = True
|
|
||||||
else:
|
|
||||||
logging.info("loaded completely {} {} {}".format(lowvram_model_memory / (1024 * 1024), mem_counter / (1024 * 1024), full_load))
|
|
||||||
self.model.model_lowvram = False
|
|
||||||
if full_load:
|
|
||||||
self.model.to(device_to)
|
|
||||||
mem_counter = self.model_size()
|
|
||||||
|
|
||||||
self.model.lowvram_patch_counter += patch_counter
|
|
||||||
self.model.device = device_to
|
|
||||||
self.model.model_loaded_weight_memory = mem_counter
|
|
||||||
|
|
||||||
def patch_model(self, device_to=None, lowvram_model_memory=0, load_weights=True, force_patch_weights=False):
|
def patch_model(self, device_to=None, lowvram_model_memory=0, load_weights=True, force_patch_weights=False):
|
||||||
for k in self.object_patches:
|
with self.use_ejected():
|
||||||
old = comfy.utils.set_attr(self.model, k, self.object_patches[k])
|
for k in self.object_patches:
|
||||||
if k not in self.object_patches_backup:
|
old = comfy.utils.set_attr(self.model, k, self.object_patches[k])
|
||||||
self.object_patches_backup[k] = old
|
if k not in self.object_patches_backup:
|
||||||
|
self.object_patches_backup[k] = old
|
||||||
|
|
||||||
if lowvram_model_memory == 0:
|
if lowvram_model_memory == 0:
|
||||||
full_load = True
|
full_load = True
|
||||||
else:
|
else:
|
||||||
full_load = False
|
full_load = False
|
||||||
|
|
||||||
if load_weights:
|
if load_weights:
|
||||||
self.load(device_to, lowvram_model_memory=lowvram_model_memory, force_patch_weights=force_patch_weights, full_load=full_load)
|
self.load(device_to, lowvram_model_memory=lowvram_model_memory, force_patch_weights=force_patch_weights, full_load=full_load)
|
||||||
|
self.inject_model()
|
||||||
return self.model
|
return self.model
|
||||||
|
|
||||||
def unpatch_model(self, device_to=None, unpatch_weights=True):
|
def unpatch_model(self, device_to=None, unpatch_weights=True):
|
||||||
|
self.eject_model()
|
||||||
if unpatch_weights:
|
if unpatch_weights:
|
||||||
|
self.unpatch_hooks()
|
||||||
if self.model.model_lowvram:
|
if self.model.model_lowvram:
|
||||||
for m in self.model.modules():
|
for m in self.model.modules():
|
||||||
wipe_lowvram_weight(m)
|
wipe_lowvram_weight(m)
|
||||||
@ -498,6 +663,7 @@ class ModelPatcher:
|
|||||||
else:
|
else:
|
||||||
comfy.utils.set_attr_param(self.model, k, bk.weight)
|
comfy.utils.set_attr_param(self.model, k, bk.weight)
|
||||||
|
|
||||||
|
self.model.current_weight_patches_uuid = None
|
||||||
self.backup.clear()
|
self.backup.clear()
|
||||||
|
|
||||||
if device_to is not None:
|
if device_to is not None:
|
||||||
@ -516,69 +682,92 @@ class ModelPatcher:
|
|||||||
self.object_patches_backup.clear()
|
self.object_patches_backup.clear()
|
||||||
|
|
||||||
def partially_unload(self, device_to, memory_to_free=0):
|
def partially_unload(self, device_to, memory_to_free=0):
|
||||||
memory_freed = 0
|
with self.use_ejected():
|
||||||
patch_counter = 0
|
memory_freed = 0
|
||||||
unload_list = self._load_list()
|
patch_counter = 0
|
||||||
unload_list.sort()
|
unload_list = self._load_list()
|
||||||
for unload in unload_list:
|
unload_list.sort()
|
||||||
if memory_to_free < memory_freed:
|
for unload in unload_list:
|
||||||
break
|
if memory_to_free < memory_freed:
|
||||||
module_mem = unload[0]
|
break
|
||||||
n = unload[1]
|
module_mem = unload[0]
|
||||||
m = unload[2]
|
n = unload[1]
|
||||||
params = unload[3]
|
m = unload[2]
|
||||||
|
params = unload[3]
|
||||||
|
|
||||||
lowvram_possible = hasattr(m, "comfy_cast_weights")
|
lowvram_possible = hasattr(m, "comfy_cast_weights")
|
||||||
if hasattr(m, "comfy_patched_weights") and m.comfy_patched_weights == True:
|
if hasattr(m, "comfy_patched_weights") and m.comfy_patched_weights == True:
|
||||||
move_weight = True
|
move_weight = True
|
||||||
for param in params:
|
for param in params:
|
||||||
key = "{}.{}".format(n, param)
|
key = "{}.{}".format(n, param)
|
||||||
bk = self.backup.get(key, None)
|
bk = self.backup.get(key, None)
|
||||||
if bk is not None:
|
if bk is not None:
|
||||||
if not lowvram_possible:
|
if not lowvram_possible:
|
||||||
move_weight = False
|
move_weight = False
|
||||||
break
|
break
|
||||||
|
|
||||||
if bk.inplace_update:
|
if bk.inplace_update:
|
||||||
comfy.utils.copy_to_param(self.model, key, bk.weight)
|
comfy.utils.copy_to_param(self.model, key, bk.weight)
|
||||||
else:
|
else:
|
||||||
comfy.utils.set_attr_param(self.model, key, bk.weight)
|
comfy.utils.set_attr_param(self.model, key, bk.weight)
|
||||||
self.backup.pop(key)
|
self.backup.pop(key)
|
||||||
|
|
||||||
|
weight_key = "{}.weight".format(n)
|
||||||
|
bias_key = "{}.bias".format(n)
|
||||||
|
if move_weight:
|
||||||
|
m.to(device_to)
|
||||||
|
if lowvram_possible:
|
||||||
|
if weight_key in self.patches:
|
||||||
|
m.weight_function = LowVramPatch(weight_key, self.patches)
|
||||||
|
patch_counter += 1
|
||||||
|
if bias_key in self.patches:
|
||||||
|
m.bias_function = LowVramPatch(bias_key, self.patches)
|
||||||
|
patch_counter += 1
|
||||||
|
|
||||||
weight_key = "{}.weight".format(n)
|
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
||||||
bias_key = "{}.bias".format(n)
|
m.comfy_cast_weights = True
|
||||||
if move_weight:
|
m.comfy_patched_weights = False
|
||||||
m.to(device_to)
|
memory_freed += module_mem
|
||||||
if lowvram_possible:
|
logging.debug("freed {}".format(n))
|
||||||
if weight_key in self.patches:
|
|
||||||
m.weight_function = LowVramPatch(weight_key, self.patches)
|
|
||||||
patch_counter += 1
|
|
||||||
if bias_key in self.patches:
|
|
||||||
m.bias_function = LowVramPatch(bias_key, self.patches)
|
|
||||||
patch_counter += 1
|
|
||||||
|
|
||||||
m.prev_comfy_cast_weights = m.comfy_cast_weights
|
self.model.model_lowvram = True
|
||||||
m.comfy_cast_weights = True
|
self.model.lowvram_patch_counter += patch_counter
|
||||||
m.comfy_patched_weights = False
|
self.model.model_loaded_weight_memory -= memory_freed
|
||||||
memory_freed += module_mem
|
return memory_freed
|
||||||
logging.debug("freed {}".format(n))
|
|
||||||
|
|
||||||
self.model.model_lowvram = True
|
def partially_load(self, device_to, extra_memory=0, force_patch_weights=False):
|
||||||
self.model.lowvram_patch_counter += patch_counter
|
with self.use_ejected(skip_and_inject_on_exit_only=True):
|
||||||
self.model.model_loaded_weight_memory -= memory_freed
|
unpatch_weights = self.model.current_weight_patches_uuid is not None and (self.model.current_weight_patches_uuid != self.patches_uuid or force_patch_weights)
|
||||||
return memory_freed
|
# TODO: force_patch_weights should not unload + reload full model
|
||||||
|
used = self.model.model_loaded_weight_memory
|
||||||
|
self.unpatch_model(self.offload_device, unpatch_weights=unpatch_weights)
|
||||||
|
if unpatch_weights:
|
||||||
|
extra_memory += (used - self.model.model_loaded_weight_memory)
|
||||||
|
|
||||||
def partially_load(self, device_to, extra_memory=0):
|
self.patch_model(load_weights=False)
|
||||||
self.unpatch_model(unpatch_weights=False)
|
full_load = False
|
||||||
self.patch_model(load_weights=False)
|
if self.model.model_lowvram == False and self.model.model_loaded_weight_memory > 0:
|
||||||
full_load = False
|
self.apply_hooks(self.forced_hooks, force_apply=True)
|
||||||
if self.model.model_lowvram == False:
|
return 0
|
||||||
return 0
|
if self.model.model_loaded_weight_memory + extra_memory > self.model_size():
|
||||||
if self.model.model_loaded_weight_memory + extra_memory > self.model_size():
|
full_load = True
|
||||||
full_load = True
|
current_used = self.model.model_loaded_weight_memory
|
||||||
current_used = self.model.model_loaded_weight_memory
|
try:
|
||||||
self.load(device_to, lowvram_model_memory=current_used + extra_memory, full_load=full_load)
|
self.load(device_to, lowvram_model_memory=current_used + extra_memory, force_patch_weights=force_patch_weights, full_load=full_load)
|
||||||
return self.model.model_loaded_weight_memory - current_used
|
except Exception as e:
|
||||||
|
self.detach()
|
||||||
|
raise e
|
||||||
|
|
||||||
|
return self.model.model_loaded_weight_memory - current_used
|
||||||
|
|
||||||
|
def detach(self, unpatch_all=True):
|
||||||
|
self.eject_model()
|
||||||
|
self.model_patches_to(self.offload_device)
|
||||||
|
if unpatch_all:
|
||||||
|
self.unpatch_model(self.offload_device, unpatch_weights=unpatch_all)
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_DETACH):
|
||||||
|
callback(self, unpatch_all)
|
||||||
|
return self.model
|
||||||
|
|
||||||
def current_loaded_device(self):
|
def current_loaded_device(self):
|
||||||
return self.model.device
|
return self.model.device
|
||||||
@ -586,3 +775,346 @@ class ModelPatcher:
|
|||||||
def calculate_weight(self, patches, weight, key, intermediate_dtype=torch.float32):
|
def calculate_weight(self, patches, weight, key, intermediate_dtype=torch.float32):
|
||||||
print("WARNING the ModelPatcher.calculate_weight function is deprecated, please use: comfy.lora.calculate_weight instead")
|
print("WARNING the ModelPatcher.calculate_weight function is deprecated, please use: comfy.lora.calculate_weight instead")
|
||||||
return comfy.lora.calculate_weight(patches, weight, key, intermediate_dtype=intermediate_dtype)
|
return comfy.lora.calculate_weight(patches, weight, key, intermediate_dtype=intermediate_dtype)
|
||||||
|
|
||||||
|
def cleanup(self):
|
||||||
|
self.clean_hooks()
|
||||||
|
if hasattr(self.model, "current_patcher"):
|
||||||
|
self.model.current_patcher = None
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_CLEANUP):
|
||||||
|
callback(self)
|
||||||
|
|
||||||
|
def add_callback(self, call_type: str, callback: Callable):
|
||||||
|
self.add_callback_with_key(call_type, None, callback)
|
||||||
|
|
||||||
|
def add_callback_with_key(self, call_type: str, key: str, callback: Callable):
|
||||||
|
c = self.callbacks.setdefault(call_type, {}).setdefault(key, [])
|
||||||
|
c.append(callback)
|
||||||
|
|
||||||
|
def remove_callbacks_with_key(self, call_type: str, key: str):
|
||||||
|
c = self.callbacks.get(call_type, {})
|
||||||
|
if key in c:
|
||||||
|
c.pop(key)
|
||||||
|
|
||||||
|
def get_callbacks(self, call_type: str, key: str):
|
||||||
|
return self.callbacks.get(call_type, {}).get(key, [])
|
||||||
|
|
||||||
|
def get_all_callbacks(self, call_type: str):
|
||||||
|
c_list = []
|
||||||
|
for c in self.callbacks.get(call_type, {}).values():
|
||||||
|
c_list.extend(c)
|
||||||
|
return c_list
|
||||||
|
|
||||||
|
def add_wrapper(self, wrapper_type: str, wrapper: Callable):
|
||||||
|
self.add_wrapper_with_key(wrapper_type, None, wrapper)
|
||||||
|
|
||||||
|
def add_wrapper_with_key(self, wrapper_type: str, key: str, wrapper: Callable):
|
||||||
|
w = self.wrappers.setdefault(wrapper_type, {}).setdefault(key, [])
|
||||||
|
w.append(wrapper)
|
||||||
|
|
||||||
|
def remove_wrappers_with_key(self, wrapper_type: str, key: str):
|
||||||
|
w = self.wrappers.get(wrapper_type, {})
|
||||||
|
if key in w:
|
||||||
|
w.pop(key)
|
||||||
|
|
||||||
|
def get_wrappers(self, wrapper_type: str, key: str):
|
||||||
|
return self.wrappers.get(wrapper_type, {}).get(key, [])
|
||||||
|
|
||||||
|
def get_all_wrappers(self, wrapper_type: str):
|
||||||
|
w_list = []
|
||||||
|
for w in self.wrappers.get(wrapper_type, {}).values():
|
||||||
|
w_list.extend(w)
|
||||||
|
return w_list
|
||||||
|
|
||||||
|
def set_attachments(self, key: str, attachment):
|
||||||
|
self.attachments[key] = attachment
|
||||||
|
|
||||||
|
def remove_attachments(self, key: str):
|
||||||
|
if key in self.attachments:
|
||||||
|
self.attachments.pop(key)
|
||||||
|
|
||||||
|
def get_attachment(self, key: str):
|
||||||
|
return self.attachments.get(key, None)
|
||||||
|
|
||||||
|
def set_injections(self, key: str, injections: list[PatcherInjection]):
|
||||||
|
self.injections[key] = injections
|
||||||
|
|
||||||
|
def remove_injections(self, key: str):
|
||||||
|
if key in self.injections:
|
||||||
|
self.injections.pop(key)
|
||||||
|
|
||||||
|
def set_additional_models(self, key: str, models: list['ModelPatcher']):
|
||||||
|
self.additional_models[key] = models
|
||||||
|
|
||||||
|
def remove_additional_models(self, key: str):
|
||||||
|
if key in self.additional_models:
|
||||||
|
self.additional_models.pop(key)
|
||||||
|
|
||||||
|
def get_additional_models_with_key(self, key: str):
|
||||||
|
return self.additional_models.get(key, [])
|
||||||
|
|
||||||
|
def get_additional_models(self):
|
||||||
|
all_models = []
|
||||||
|
for models in self.additional_models.values():
|
||||||
|
all_models.extend(models)
|
||||||
|
return all_models
|
||||||
|
|
||||||
|
def get_nested_additional_models(self):
|
||||||
|
def _evaluate_sub_additional_models(prev_models: list[ModelPatcher], cache_set: set[ModelPatcher]):
|
||||||
|
'''Make sure circular references do not cause infinite recursion.'''
|
||||||
|
next_models = []
|
||||||
|
for model in prev_models:
|
||||||
|
candidates = model.get_additional_models()
|
||||||
|
for c in candidates:
|
||||||
|
if c not in cache_set:
|
||||||
|
next_models.append(c)
|
||||||
|
cache_set.add(c)
|
||||||
|
if len(next_models) == 0:
|
||||||
|
return prev_models
|
||||||
|
return prev_models + _evaluate_sub_additional_models(next_models, cache_set)
|
||||||
|
|
||||||
|
all_models = self.get_additional_models()
|
||||||
|
models_set = set(all_models)
|
||||||
|
real_all_models = _evaluate_sub_additional_models(prev_models=all_models, cache_set=models_set)
|
||||||
|
return real_all_models
|
||||||
|
|
||||||
|
def use_ejected(self, skip_and_inject_on_exit_only=False):
|
||||||
|
return AutoPatcherEjector(self, skip_and_inject_on_exit_only=skip_and_inject_on_exit_only)
|
||||||
|
|
||||||
|
def inject_model(self):
|
||||||
|
if self.is_injected or self.skip_injection:
|
||||||
|
return
|
||||||
|
for injections in self.injections.values():
|
||||||
|
for inj in injections:
|
||||||
|
inj.inject(self)
|
||||||
|
self.is_injected = True
|
||||||
|
if self.is_injected:
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_INJECT_MODEL):
|
||||||
|
callback(self)
|
||||||
|
|
||||||
|
def eject_model(self):
|
||||||
|
if not self.is_injected:
|
||||||
|
return
|
||||||
|
for injections in self.injections.values():
|
||||||
|
for inj in injections:
|
||||||
|
inj.eject(self)
|
||||||
|
self.is_injected = False
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_EJECT_MODEL):
|
||||||
|
callback(self)
|
||||||
|
|
||||||
|
def pre_run(self):
|
||||||
|
if hasattr(self.model, "current_patcher"):
|
||||||
|
self.model.current_patcher = self
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_PRE_RUN):
|
||||||
|
callback(self)
|
||||||
|
|
||||||
|
def prepare_state(self, timestep):
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_PREPARE_STATE):
|
||||||
|
callback(self, timestep)
|
||||||
|
|
||||||
|
def restore_hook_patches(self):
|
||||||
|
if len(self.hook_patches_backup) > 0:
|
||||||
|
self.hook_patches = self.hook_patches_backup
|
||||||
|
self.hook_patches_backup = {}
|
||||||
|
|
||||||
|
def set_hook_mode(self, hook_mode: comfy.hooks.EnumHookMode):
|
||||||
|
self.hook_mode = hook_mode
|
||||||
|
|
||||||
|
def prepare_hook_patches_current_keyframe(self, t: torch.Tensor, hook_group: comfy.hooks.HookGroup):
|
||||||
|
curr_t = t[0]
|
||||||
|
reset_current_hooks = False
|
||||||
|
for hook in hook_group.hooks:
|
||||||
|
changed = hook.hook_keyframe.prepare_current_keyframe(curr_t=curr_t)
|
||||||
|
# if keyframe changed, remove any cached HookGroups that contain hook with the same hook_ref;
|
||||||
|
# this will cause the weights to be recalculated when sampling
|
||||||
|
if changed:
|
||||||
|
# reset current_hooks if contains hook that changed
|
||||||
|
if self.current_hooks is not None:
|
||||||
|
for current_hook in self.current_hooks.hooks:
|
||||||
|
if current_hook == hook:
|
||||||
|
reset_current_hooks = True
|
||||||
|
break
|
||||||
|
for cached_group in list(self.cached_hook_patches.keys()):
|
||||||
|
if cached_group.contains(hook):
|
||||||
|
self.cached_hook_patches.pop(cached_group)
|
||||||
|
if reset_current_hooks:
|
||||||
|
self.patch_hooks(None)
|
||||||
|
|
||||||
|
def register_all_hook_patches(self, hooks_dict: dict[comfy.hooks.EnumHookType, dict[comfy.hooks.Hook, None]], target: comfy.hooks.EnumWeightTarget, model_options: dict=None):
|
||||||
|
self.restore_hook_patches()
|
||||||
|
registered_hooks: list[comfy.hooks.Hook] = []
|
||||||
|
# handle WrapperHooks, if model_options provided
|
||||||
|
if model_options is not None:
|
||||||
|
for hook in hooks_dict.get(comfy.hooks.EnumHookType.Wrappers, {}):
|
||||||
|
hook.add_hook_patches(self, model_options, target, registered_hooks)
|
||||||
|
# handle WeightHooks
|
||||||
|
weight_hooks_to_register: list[comfy.hooks.WeightHook] = []
|
||||||
|
for hook in hooks_dict.get(comfy.hooks.EnumHookType.Weight, {}):
|
||||||
|
if hook.hook_ref not in self.hook_patches:
|
||||||
|
weight_hooks_to_register.append(hook)
|
||||||
|
if len(weight_hooks_to_register) > 0:
|
||||||
|
# clone hook_patches to become backup so that any non-dynamic hooks will return to their original state
|
||||||
|
self.hook_patches_backup = create_hook_patches_clone(self.hook_patches)
|
||||||
|
for hook in weight_hooks_to_register:
|
||||||
|
hook.add_hook_patches(self, model_options, target, registered_hooks)
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_REGISTER_ALL_HOOK_PATCHES):
|
||||||
|
callback(self, hooks_dict, target)
|
||||||
|
|
||||||
|
def add_hook_patches(self, hook: comfy.hooks.WeightHook, patches, strength_patch=1.0, strength_model=1.0):
|
||||||
|
with self.use_ejected():
|
||||||
|
# NOTE: this mirrors behavior of add_patches func
|
||||||
|
current_hook_patches: dict[str,list] = self.hook_patches.get(hook.hook_ref, {})
|
||||||
|
p = set()
|
||||||
|
model_sd = self.model.state_dict()
|
||||||
|
for k in patches:
|
||||||
|
offset = None
|
||||||
|
function = None
|
||||||
|
if isinstance(k, str):
|
||||||
|
key = k
|
||||||
|
else:
|
||||||
|
offset = k[1]
|
||||||
|
key = k[0]
|
||||||
|
if len(k) > 2:
|
||||||
|
function = k[2]
|
||||||
|
|
||||||
|
if key in model_sd:
|
||||||
|
p.add(k)
|
||||||
|
current_patches: list[tuple] = current_hook_patches.get(key, [])
|
||||||
|
current_patches.append((strength_patch, patches[k], strength_model, offset, function))
|
||||||
|
current_hook_patches[key] = current_patches
|
||||||
|
self.hook_patches[hook.hook_ref] = current_hook_patches
|
||||||
|
# since should care about these patches too to determine if same model, reroll patches_uuid
|
||||||
|
self.patches_uuid = uuid.uuid4()
|
||||||
|
return list(p)
|
||||||
|
|
||||||
|
def get_combined_hook_patches(self, hooks: comfy.hooks.HookGroup):
|
||||||
|
# combined_patches will contain weights of all relevant hooks, per key
|
||||||
|
combined_patches = {}
|
||||||
|
if hooks is not None:
|
||||||
|
for hook in hooks.hooks:
|
||||||
|
hook_patches: dict = self.hook_patches.get(hook.hook_ref, {})
|
||||||
|
for key in hook_patches.keys():
|
||||||
|
current_patches: list[tuple] = combined_patches.get(key, [])
|
||||||
|
if math.isclose(hook.strength, 1.0):
|
||||||
|
current_patches.extend(hook_patches[key])
|
||||||
|
else:
|
||||||
|
# patches are stored as tuples: (strength_patch, (tuple_with_weights,), strength_model)
|
||||||
|
for patch in hook_patches[key]:
|
||||||
|
new_patch = list(patch)
|
||||||
|
new_patch[0] *= hook.strength
|
||||||
|
current_patches.append(tuple(new_patch))
|
||||||
|
combined_patches[key] = current_patches
|
||||||
|
return combined_patches
|
||||||
|
|
||||||
|
def apply_hooks(self, hooks: comfy.hooks.HookGroup, transformer_options: dict=None, force_apply=False):
|
||||||
|
# TODO: return transformer_options dict with any additions from hooks
|
||||||
|
if self.current_hooks == hooks and (not force_apply or (not self.is_clip and hooks is None)):
|
||||||
|
return {}
|
||||||
|
self.patch_hooks(hooks=hooks)
|
||||||
|
for callback in self.get_all_callbacks(CallbacksMP.ON_APPLY_HOOKS):
|
||||||
|
callback(self, hooks)
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def patch_hooks(self, hooks: comfy.hooks.HookGroup):
|
||||||
|
with self.use_ejected():
|
||||||
|
self.unpatch_hooks()
|
||||||
|
if hooks is not None:
|
||||||
|
model_sd_keys = list(self.model_state_dict().keys())
|
||||||
|
memory_counter = None
|
||||||
|
if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed:
|
||||||
|
# TODO: minimum_counter should have a minimum that conforms to loaded model requirements
|
||||||
|
memory_counter = MemoryCounter(initial=comfy.model_management.get_free_memory(self.load_device),
|
||||||
|
minimum=comfy.model_management.minimum_inference_memory()*2)
|
||||||
|
# if have cached weights for hooks, use it
|
||||||
|
cached_weights = self.cached_hook_patches.get(hooks, None)
|
||||||
|
if cached_weights is not None:
|
||||||
|
for key in cached_weights:
|
||||||
|
if key not in model_sd_keys:
|
||||||
|
print(f"WARNING cached hook could not patch. key does not exist in model: {key}")
|
||||||
|
continue
|
||||||
|
self.patch_cached_hook_weights(cached_weights=cached_weights, key=key, memory_counter=memory_counter)
|
||||||
|
else:
|
||||||
|
relevant_patches = self.get_combined_hook_patches(hooks=hooks)
|
||||||
|
original_weights = None
|
||||||
|
if len(relevant_patches) > 0:
|
||||||
|
original_weights = self.get_key_patches()
|
||||||
|
for key in relevant_patches:
|
||||||
|
if key not in model_sd_keys:
|
||||||
|
print(f"WARNING cached hook would not patch. key does not exist in model: {key}")
|
||||||
|
continue
|
||||||
|
self.patch_hook_weight_to_device(hooks=hooks, combined_patches=relevant_patches, key=key, original_weights=original_weights,
|
||||||
|
memory_counter=memory_counter)
|
||||||
|
self.current_hooks = hooks
|
||||||
|
|
||||||
|
def patch_cached_hook_weights(self, cached_weights: dict, key: str, memory_counter: MemoryCounter):
|
||||||
|
if key not in self.hook_backup:
|
||||||
|
weight: torch.Tensor = comfy.utils.get_attr(self.model, key)
|
||||||
|
target_device = self.offload_device
|
||||||
|
if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed:
|
||||||
|
used = memory_counter.use(weight)
|
||||||
|
if used:
|
||||||
|
target_device = weight.device
|
||||||
|
self.hook_backup[key] = (weight.to(device=target_device, copy=True), weight.device)
|
||||||
|
comfy.utils.copy_to_param(self.model, key, cached_weights[key][0].to(device=cached_weights[key][1]))
|
||||||
|
|
||||||
|
def clear_cached_hook_weights(self):
|
||||||
|
self.cached_hook_patches.clear()
|
||||||
|
self.patch_hooks(None)
|
||||||
|
|
||||||
|
def patch_hook_weight_to_device(self, hooks: comfy.hooks.HookGroup, combined_patches: dict, key: str, original_weights: dict, memory_counter: MemoryCounter):
|
||||||
|
if key not in combined_patches:
|
||||||
|
return
|
||||||
|
|
||||||
|
weight, set_func, convert_func = get_key_weight(self.model, key)
|
||||||
|
weight: torch.Tensor
|
||||||
|
if key not in self.hook_backup:
|
||||||
|
target_device = self.offload_device
|
||||||
|
if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed:
|
||||||
|
used = memory_counter.use(weight)
|
||||||
|
if used:
|
||||||
|
target_device = weight.device
|
||||||
|
self.hook_backup[key] = (weight.to(device=target_device, copy=True), weight.device)
|
||||||
|
# TODO: properly handle LowVramPatch, if it ends up an issue
|
||||||
|
temp_weight = comfy.model_management.cast_to_device(weight, weight.device, torch.float32, copy=True)
|
||||||
|
if convert_func is not None:
|
||||||
|
temp_weight = convert_func(temp_weight, inplace=True)
|
||||||
|
|
||||||
|
out_weight = comfy.lora.calculate_weight(combined_patches[key],
|
||||||
|
temp_weight,
|
||||||
|
key, original_weights=original_weights)
|
||||||
|
del original_weights[key]
|
||||||
|
if set_func is None:
|
||||||
|
out_weight = comfy.float.stochastic_rounding(out_weight, weight.dtype, seed=string_to_seed(key))
|
||||||
|
comfy.utils.copy_to_param(self.model, key, out_weight)
|
||||||
|
else:
|
||||||
|
set_func(out_weight, inplace_update=True, seed=string_to_seed(key))
|
||||||
|
if self.hook_mode == comfy.hooks.EnumHookMode.MaxSpeed:
|
||||||
|
# TODO: disable caching if not enough system RAM to do so
|
||||||
|
target_device = self.offload_device
|
||||||
|
used = memory_counter.use(weight)
|
||||||
|
if used:
|
||||||
|
target_device = weight.device
|
||||||
|
self.cached_hook_patches.setdefault(hooks, {})
|
||||||
|
self.cached_hook_patches[hooks][key] = (out_weight.to(device=target_device, copy=False), weight.device)
|
||||||
|
del temp_weight
|
||||||
|
del out_weight
|
||||||
|
del weight
|
||||||
|
|
||||||
|
def unpatch_hooks(self) -> None:
|
||||||
|
with self.use_ejected():
|
||||||
|
if len(self.hook_backup) == 0:
|
||||||
|
self.current_hooks = None
|
||||||
|
return
|
||||||
|
keys = list(self.hook_backup.keys())
|
||||||
|
for k in keys:
|
||||||
|
comfy.utils.copy_to_param(self.model, k, self.hook_backup[k][0].to(device=self.hook_backup[k][1]))
|
||||||
|
|
||||||
|
self.hook_backup.clear()
|
||||||
|
self.current_hooks = None
|
||||||
|
|
||||||
|
def clean_hooks(self):
|
||||||
|
self.unpatch_hooks()
|
||||||
|
self.clear_cached_hook_weights()
|
||||||
|
|
||||||
|
def __del__(self):
|
||||||
|
self.detach(unpatch_all=False)
|
||||||
|
|
||||||
|
|||||||
@ -243,7 +243,7 @@ class ModelSamplingDiscreteFlow(torch.nn.Module):
|
|||||||
return 1.0
|
return 1.0
|
||||||
if percent >= 1.0:
|
if percent >= 1.0:
|
||||||
return 0.0
|
return 0.0
|
||||||
return 1.0 - percent
|
return time_snr_shift(self.shift, 1.0 - percent)
|
||||||
|
|
||||||
class StableCascadeSampling(ModelSamplingDiscrete):
|
class StableCascadeSampling(ModelSamplingDiscrete):
|
||||||
def __init__(self, model_config=None):
|
def __init__(self, model_config=None):
|
||||||
@ -336,4 +336,4 @@ class ModelSamplingFlux(torch.nn.Module):
|
|||||||
return 1.0
|
return 1.0
|
||||||
if percent >= 1.0:
|
if percent >= 1.0:
|
||||||
return 0.0
|
return 0.0
|
||||||
return 1.0 - percent
|
return flux_time_shift(self.shift, 1.0, 1.0 - percent)
|
||||||
|
|||||||
@ -269,7 +269,7 @@ def fp8_linear(self, input):
|
|||||||
|
|
||||||
if scale_input is None:
|
if scale_input is None:
|
||||||
scale_input = torch.ones((), device=input.device, dtype=torch.float32)
|
scale_input = torch.ones((), device=input.device, dtype=torch.float32)
|
||||||
inn = input.reshape(-1, input.shape[2]).to(dtype)
|
inn = torch.clamp(input, min=-448, max=448).reshape(-1, input.shape[2]).to(dtype)
|
||||||
else:
|
else:
|
||||||
scale_input = scale_input.to(input.device)
|
scale_input = scale_input.to(input.device)
|
||||||
inn = (input * (1.0 / scale_input).to(input.dtype)).reshape(-1, input.shape[2]).to(dtype)
|
inn = (input * (1.0 / scale_input).to(input.dtype)).reshape(-1, input.shape[2]).to(dtype)
|
||||||
|
|||||||
156
comfy/patcher_extension.py
Normal file
156
comfy/patcher_extension.py
Normal file
@ -0,0 +1,156 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from typing import Callable
|
||||||
|
|
||||||
|
class CallbacksMP:
|
||||||
|
ON_CLONE = "on_clone"
|
||||||
|
ON_LOAD = "on_load_after"
|
||||||
|
ON_DETACH = "on_detach_after"
|
||||||
|
ON_CLEANUP = "on_cleanup"
|
||||||
|
ON_PRE_RUN = "on_pre_run"
|
||||||
|
ON_PREPARE_STATE = "on_prepare_state"
|
||||||
|
ON_APPLY_HOOKS = "on_apply_hooks"
|
||||||
|
ON_REGISTER_ALL_HOOK_PATCHES = "on_register_all_hook_patches"
|
||||||
|
ON_INJECT_MODEL = "on_inject_model"
|
||||||
|
ON_EJECT_MODEL = "on_eject_model"
|
||||||
|
|
||||||
|
# callbacks dict is in the format:
|
||||||
|
# {"call_type": {"key": [Callable1, Callable2, ...]} }
|
||||||
|
@classmethod
|
||||||
|
def init_callbacks(cls) -> dict[str, dict[str, list[Callable]]]:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def add_callback(call_type: str, callback: Callable, transformer_options: dict, is_model_options=False):
|
||||||
|
add_callback_with_key(call_type, None, callback, transformer_options, is_model_options)
|
||||||
|
|
||||||
|
def add_callback_with_key(call_type: str, key: str, callback: Callable, transformer_options: dict, is_model_options=False):
|
||||||
|
if is_model_options:
|
||||||
|
transformer_options = transformer_options.setdefault("transformer_options", {})
|
||||||
|
callbacks: dict[str, dict[str, list]] = transformer_options.setdefault("callbacks", {})
|
||||||
|
c = callbacks.setdefault(call_type, {}).setdefault(key, [])
|
||||||
|
c.append(callback)
|
||||||
|
|
||||||
|
def get_callbacks_with_key(call_type: str, key: str, transformer_options: dict, is_model_options=False):
|
||||||
|
if is_model_options:
|
||||||
|
transformer_options = transformer_options.get("transformer_options", {})
|
||||||
|
c_list = []
|
||||||
|
callbacks: dict[str, list] = transformer_options.get("callbacks", {})
|
||||||
|
c_list.extend(callbacks.get(call_type, {}).get(key, []))
|
||||||
|
return c_list
|
||||||
|
|
||||||
|
def get_all_callbacks(call_type: str, transformer_options: dict, is_model_options=False):
|
||||||
|
if is_model_options:
|
||||||
|
transformer_options = transformer_options.get("transformer_options", {})
|
||||||
|
c_list = []
|
||||||
|
callbacks: dict[str, list] = transformer_options.get("callbacks", {})
|
||||||
|
for c in callbacks.get(call_type, {}).values():
|
||||||
|
c_list.extend(c)
|
||||||
|
return c_list
|
||||||
|
|
||||||
|
class WrappersMP:
|
||||||
|
OUTER_SAMPLE = "outer_sample"
|
||||||
|
SAMPLER_SAMPLE = "sampler_sample"
|
||||||
|
CALC_COND_BATCH = "calc_cond_batch"
|
||||||
|
APPLY_MODEL = "apply_model"
|
||||||
|
DIFFUSION_MODEL = "diffusion_model"
|
||||||
|
|
||||||
|
# wrappers dict is in the format:
|
||||||
|
# {"wrapper_type": {"key": [Callable1, Callable2, ...]} }
|
||||||
|
@classmethod
|
||||||
|
def init_wrappers(cls) -> dict[str, dict[str, list[Callable]]]:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
def add_wrapper(wrapper_type: str, wrapper: Callable, transformer_options: dict, is_model_options=False):
|
||||||
|
add_wrapper_with_key(wrapper_type, None, wrapper, transformer_options, is_model_options)
|
||||||
|
|
||||||
|
def add_wrapper_with_key(wrapper_type: str, key: str, wrapper: Callable, transformer_options: dict, is_model_options=False):
|
||||||
|
if is_model_options:
|
||||||
|
transformer_options = transformer_options.setdefault("transformer_options", {})
|
||||||
|
wrappers: dict[str, dict[str, list]] = transformer_options.setdefault("wrappers", {})
|
||||||
|
w = wrappers.setdefault(wrapper_type, {}).setdefault(key, [])
|
||||||
|
w.append(wrapper)
|
||||||
|
|
||||||
|
def get_wrappers_with_key(wrapper_type: str, key: str, transformer_options: dict, is_model_options=False):
|
||||||
|
if is_model_options:
|
||||||
|
transformer_options = transformer_options.get("transformer_options", {})
|
||||||
|
w_list = []
|
||||||
|
wrappers: dict[str, list] = transformer_options.get("wrappers", {})
|
||||||
|
w_list.extend(wrappers.get(wrapper_type, {}).get(key, []))
|
||||||
|
return w_list
|
||||||
|
|
||||||
|
def get_all_wrappers(wrapper_type: str, transformer_options: dict, is_model_options=False):
|
||||||
|
if is_model_options:
|
||||||
|
transformer_options = transformer_options.get("transformer_options", {})
|
||||||
|
w_list = []
|
||||||
|
wrappers: dict[str, list] = transformer_options.get("wrappers", {})
|
||||||
|
for w in wrappers.get(wrapper_type, {}).values():
|
||||||
|
w_list.extend(w)
|
||||||
|
return w_list
|
||||||
|
|
||||||
|
class WrapperExecutor:
|
||||||
|
"""Handles call stack of wrappers around a function in an ordered manner."""
|
||||||
|
def __init__(self, original: Callable, class_obj: object, wrappers: list[Callable], idx: int):
|
||||||
|
# NOTE: class_obj exists so that wrappers surrounding a class method can access
|
||||||
|
# the class instance at runtime via executor.class_obj
|
||||||
|
self.original = original
|
||||||
|
self.class_obj = class_obj
|
||||||
|
self.wrappers = wrappers.copy()
|
||||||
|
self.idx = idx
|
||||||
|
self.is_last = idx == len(wrappers)
|
||||||
|
|
||||||
|
def __call__(self, *args, **kwargs):
|
||||||
|
"""Calls the next wrapper or original function, whichever is appropriate."""
|
||||||
|
new_executor = self._create_next_executor()
|
||||||
|
return new_executor.execute(*args, **kwargs)
|
||||||
|
|
||||||
|
def execute(self, *args, **kwargs):
|
||||||
|
"""Used to initiate executor internally - DO NOT use this if you received executor in wrapper."""
|
||||||
|
args = list(args)
|
||||||
|
kwargs = dict(kwargs)
|
||||||
|
if self.is_last:
|
||||||
|
return self.original(*args, **kwargs)
|
||||||
|
return self.wrappers[self.idx](self, *args, **kwargs)
|
||||||
|
|
||||||
|
def _create_next_executor(self) -> 'WrapperExecutor':
|
||||||
|
new_idx = self.idx + 1
|
||||||
|
if new_idx > len(self.wrappers):
|
||||||
|
raise Exception(f"Wrapper idx exceeded available wrappers; something went very wrong.")
|
||||||
|
if self.class_obj is None:
|
||||||
|
return WrapperExecutor.new_executor(self.original, self.wrappers, new_idx)
|
||||||
|
return WrapperExecutor.new_class_executor(self.original, self.class_obj, self.wrappers, new_idx)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def new_executor(cls, original: Callable, wrappers: list[Callable], idx=0):
|
||||||
|
return cls(original, class_obj=None, wrappers=wrappers, idx=idx)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def new_class_executor(cls, original: Callable, class_obj: object, wrappers: list[Callable], idx=0):
|
||||||
|
return cls(original, class_obj, wrappers, idx=idx)
|
||||||
|
|
||||||
|
class PatcherInjection:
|
||||||
|
def __init__(self, inject: Callable, eject: Callable):
|
||||||
|
self.inject = inject
|
||||||
|
self.eject = eject
|
||||||
|
|
||||||
|
def copy_nested_dicts(input_dict: dict):
|
||||||
|
new_dict = input_dict.copy()
|
||||||
|
for key, value in input_dict.items():
|
||||||
|
if isinstance(value, dict):
|
||||||
|
new_dict[key] = copy_nested_dicts(value)
|
||||||
|
elif isinstance(value, list):
|
||||||
|
new_dict[key] = value.copy()
|
||||||
|
return new_dict
|
||||||
|
|
||||||
|
def merge_nested_dicts(dict1: dict, dict2: dict, copy_dict1=True):
|
||||||
|
if copy_dict1:
|
||||||
|
merged_dict = copy_nested_dicts(dict1)
|
||||||
|
else:
|
||||||
|
merged_dict = dict1
|
||||||
|
for key, value in dict2.items():
|
||||||
|
if isinstance(value, dict):
|
||||||
|
curr_value = merged_dict.setdefault(key, {})
|
||||||
|
merged_dict[key] = merge_nested_dicts(value, curr_value)
|
||||||
|
elif isinstance(value, list):
|
||||||
|
merged_dict.setdefault(key, []).extend(value)
|
||||||
|
else:
|
||||||
|
merged_dict[key] = value
|
||||||
|
return merged_dict
|
||||||
@ -1,7 +1,15 @@
|
|||||||
import torch
|
from __future__ import annotations
|
||||||
|
import uuid
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
import comfy.conds
|
import comfy.conds
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
|
import comfy.hooks
|
||||||
|
import comfy.patcher_extension
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from comfy.model_patcher import ModelPatcher
|
||||||
|
from comfy.model_base import BaseModel
|
||||||
|
from comfy.controlnet import ControlBase
|
||||||
|
|
||||||
def prepare_mask(noise_mask, shape, device):
|
def prepare_mask(noise_mask, shape, device):
|
||||||
return comfy.utils.reshape_mask(noise_mask, shape).to(device)
|
return comfy.utils.reshape_mask(noise_mask, shape).to(device)
|
||||||
@ -10,9 +18,43 @@ def get_models_from_cond(cond, model_type):
|
|||||||
models = []
|
models = []
|
||||||
for c in cond:
|
for c in cond:
|
||||||
if model_type in c:
|
if model_type in c:
|
||||||
models += [c[model_type]]
|
if isinstance(c[model_type], list):
|
||||||
|
models += c[model_type]
|
||||||
|
else:
|
||||||
|
models += [c[model_type]]
|
||||||
return models
|
return models
|
||||||
|
|
||||||
|
def get_hooks_from_cond(cond, hooks_dict: dict[comfy.hooks.EnumHookType, dict[comfy.hooks.Hook, None]]):
|
||||||
|
# get hooks from conds, and collect cnets so they can be checked for extra_hooks
|
||||||
|
cnets: list[ControlBase] = []
|
||||||
|
for c in cond:
|
||||||
|
if 'hooks' in c:
|
||||||
|
for hook in c['hooks'].hooks:
|
||||||
|
hook: comfy.hooks.Hook
|
||||||
|
with_type = hooks_dict.setdefault(hook.hook_type, {})
|
||||||
|
with_type[hook] = None
|
||||||
|
if 'control' in c:
|
||||||
|
cnets.append(c['control'])
|
||||||
|
|
||||||
|
def get_extra_hooks_from_cnet(cnet: ControlBase, _list: list):
|
||||||
|
if cnet.extra_hooks is not None:
|
||||||
|
_list.append(cnet.extra_hooks)
|
||||||
|
if cnet.previous_controlnet is None:
|
||||||
|
return _list
|
||||||
|
return get_extra_hooks_from_cnet(cnet.previous_controlnet, _list)
|
||||||
|
|
||||||
|
hooks_list = []
|
||||||
|
cnets = set(cnets)
|
||||||
|
for base_cnet in cnets:
|
||||||
|
get_extra_hooks_from_cnet(base_cnet, hooks_list)
|
||||||
|
extra_hooks = comfy.hooks.HookGroup.combine_all_hooks(hooks_list)
|
||||||
|
if extra_hooks is not None:
|
||||||
|
for hook in extra_hooks.hooks:
|
||||||
|
with_type = hooks_dict.setdefault(hook.hook_type, {})
|
||||||
|
with_type[hook] = None
|
||||||
|
|
||||||
|
return hooks_dict
|
||||||
|
|
||||||
def convert_cond(cond):
|
def convert_cond(cond):
|
||||||
out = []
|
out = []
|
||||||
for c in cond:
|
for c in cond:
|
||||||
@ -22,17 +64,22 @@ def convert_cond(cond):
|
|||||||
model_conds["c_crossattn"] = comfy.conds.CONDCrossAttn(c[0]) #TODO: remove
|
model_conds["c_crossattn"] = comfy.conds.CONDCrossAttn(c[0]) #TODO: remove
|
||||||
temp["cross_attn"] = c[0]
|
temp["cross_attn"] = c[0]
|
||||||
temp["model_conds"] = model_conds
|
temp["model_conds"] = model_conds
|
||||||
|
temp["uuid"] = uuid.uuid4()
|
||||||
out.append(temp)
|
out.append(temp)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def get_additional_models(conds, dtype):
|
def get_additional_models(conds, dtype):
|
||||||
"""loads additional models in conditioning"""
|
"""loads additional models in conditioning"""
|
||||||
cnets = []
|
cnets: list[ControlBase] = []
|
||||||
gligen = []
|
gligen = []
|
||||||
|
add_models = []
|
||||||
|
hooks: dict[comfy.hooks.EnumHookType, dict[comfy.hooks.Hook, None]] = {}
|
||||||
|
|
||||||
for k in conds:
|
for k in conds:
|
||||||
cnets += get_models_from_cond(conds[k], "control")
|
cnets += get_models_from_cond(conds[k], "control")
|
||||||
gligen += get_models_from_cond(conds[k], "gligen")
|
gligen += get_models_from_cond(conds[k], "gligen")
|
||||||
|
add_models += get_models_from_cond(conds[k], "additional_models")
|
||||||
|
get_hooks_from_cond(conds[k], hooks)
|
||||||
|
|
||||||
control_nets = set(cnets)
|
control_nets = set(cnets)
|
||||||
|
|
||||||
@ -43,7 +90,9 @@ def get_additional_models(conds, dtype):
|
|||||||
inference_memory += m.inference_memory_requirements(dtype)
|
inference_memory += m.inference_memory_requirements(dtype)
|
||||||
|
|
||||||
gligen = [x[1] for x in gligen]
|
gligen = [x[1] for x in gligen]
|
||||||
models = control_models + gligen
|
hook_models = [x.model for x in hooks.get(comfy.hooks.EnumHookType.AddModels, {}).keys()]
|
||||||
|
models = control_models + gligen + add_models + hook_models
|
||||||
|
|
||||||
return models, inference_memory
|
return models, inference_memory
|
||||||
|
|
||||||
def cleanup_additional_models(models):
|
def cleanup_additional_models(models):
|
||||||
@ -53,10 +102,11 @@ def cleanup_additional_models(models):
|
|||||||
m.cleanup()
|
m.cleanup()
|
||||||
|
|
||||||
|
|
||||||
def prepare_sampling(model, noise_shape, conds):
|
def prepare_sampling(model: 'ModelPatcher', noise_shape, conds):
|
||||||
device = model.load_device
|
device = model.load_device
|
||||||
real_model = None
|
real_model: 'BaseModel' = None
|
||||||
models, inference_memory = get_additional_models(conds, model.model_dtype())
|
models, inference_memory = get_additional_models(conds, model.model_dtype())
|
||||||
|
models += model.get_nested_additional_models() # TODO: does this require inference_memory update?
|
||||||
memory_required = model.memory_required([noise_shape[0] * 2] + list(noise_shape[1:])) + inference_memory
|
memory_required = model.memory_required([noise_shape[0] * 2] + list(noise_shape[1:])) + inference_memory
|
||||||
minimum_memory_required = model.memory_required([noise_shape[0]] + list(noise_shape[1:])) + inference_memory
|
minimum_memory_required = model.memory_required([noise_shape[0]] + list(noise_shape[1:])) + inference_memory
|
||||||
comfy.model_management.load_models_gpu([model] + models, memory_required=memory_required, minimum_memory_required=minimum_memory_required)
|
comfy.model_management.load_models_gpu([model] + models, memory_required=memory_required, minimum_memory_required=minimum_memory_required)
|
||||||
@ -72,3 +122,14 @@ def cleanup_models(conds, models):
|
|||||||
control_cleanup += get_models_from_cond(conds[k], "control")
|
control_cleanup += get_models_from_cond(conds[k], "control")
|
||||||
|
|
||||||
cleanup_additional_models(set(control_cleanup))
|
cleanup_additional_models(set(control_cleanup))
|
||||||
|
|
||||||
|
def prepare_model_patcher(model: 'ModelPatcher', conds, model_options: dict):
|
||||||
|
# check for hooks in conds - if not registered, see if can be applied
|
||||||
|
hooks = {}
|
||||||
|
for k in conds:
|
||||||
|
get_hooks_from_cond(conds[k], hooks)
|
||||||
|
# add wrappers and callbacks from ModelPatcher to transformer_options
|
||||||
|
model_options["transformer_options"]["wrappers"] = comfy.patcher_extension.copy_nested_dicts(model.wrappers)
|
||||||
|
model_options["transformer_options"]["callbacks"] = comfy.patcher_extension.copy_nested_dicts(model.callbacks)
|
||||||
|
# register hooks on model/model_options
|
||||||
|
model.register_all_hook_patches(hooks, comfy.hooks.EnumWeightTarget.Model, model_options)
|
||||||
|
|||||||
@ -1,11 +1,21 @@
|
|||||||
|
from __future__ import annotations
|
||||||
from .k_diffusion import sampling as k_diffusion_sampling
|
from .k_diffusion import sampling as k_diffusion_sampling
|
||||||
from .extra_samplers import uni_pc
|
from .extra_samplers import uni_pc
|
||||||
|
from typing import TYPE_CHECKING
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from comfy.model_patcher import ModelPatcher
|
||||||
|
from comfy.model_base import BaseModel
|
||||||
|
from comfy.controlnet import ControlBase
|
||||||
import torch
|
import torch
|
||||||
import collections
|
import collections
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
import math
|
import math
|
||||||
import logging
|
import logging
|
||||||
|
import comfy.samplers
|
||||||
import comfy.sampler_helpers
|
import comfy.sampler_helpers
|
||||||
|
import comfy.model_patcher
|
||||||
|
import comfy.patcher_extension
|
||||||
|
import comfy.hooks
|
||||||
import scipy.stats
|
import scipy.stats
|
||||||
import numpy
|
import numpy
|
||||||
|
|
||||||
@ -70,6 +80,7 @@ def get_area_and_mult(conds, x_in, timestep_in):
|
|||||||
for c in model_conds:
|
for c in model_conds:
|
||||||
conditioning[c] = model_conds[c].process_cond(batch_size=x_in.shape[0], device=x_in.device, area=area)
|
conditioning[c] = model_conds[c].process_cond(batch_size=x_in.shape[0], device=x_in.device, area=area)
|
||||||
|
|
||||||
|
hooks = conds.get('hooks', None)
|
||||||
control = conds.get('control', None)
|
control = conds.get('control', None)
|
||||||
|
|
||||||
patches = None
|
patches = None
|
||||||
@ -85,8 +96,8 @@ def get_area_and_mult(conds, x_in, timestep_in):
|
|||||||
|
|
||||||
patches['middle_patch'] = [gligen_patch]
|
patches['middle_patch'] = [gligen_patch]
|
||||||
|
|
||||||
cond_obj = collections.namedtuple('cond_obj', ['input_x', 'mult', 'conditioning', 'area', 'control', 'patches'])
|
cond_obj = collections.namedtuple('cond_obj', ['input_x', 'mult', 'conditioning', 'area', 'control', 'patches', 'uuid', 'hooks'])
|
||||||
return cond_obj(input_x, mult, conditioning, area, control, patches)
|
return cond_obj(input_x, mult, conditioning, area, control, patches, conds['uuid'], hooks)
|
||||||
|
|
||||||
def cond_equal_size(c1, c2):
|
def cond_equal_size(c1, c2):
|
||||||
if c1 is c2:
|
if c1 is c2:
|
||||||
@ -138,110 +149,184 @@ def cond_cat(c_list):
|
|||||||
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def calc_cond_batch(model, conds, x_in, timestep, model_options):
|
def finalize_default_conds(model: 'BaseModel', hooked_to_run: dict[comfy.hooks.HookGroup,list[tuple[tuple,int]]], default_conds: list[list[dict]], x_in, timestep):
|
||||||
|
# need to figure out remaining unmasked area for conds
|
||||||
|
default_mults = []
|
||||||
|
for _ in default_conds:
|
||||||
|
default_mults.append(torch.ones_like(x_in))
|
||||||
|
# look through each finalized cond in hooked_to_run for 'mult' and subtract it from each cond
|
||||||
|
for lora_hooks, to_run in hooked_to_run.items():
|
||||||
|
for cond_obj, i in to_run:
|
||||||
|
# if no default_cond for cond_type, do nothing
|
||||||
|
if len(default_conds[i]) == 0:
|
||||||
|
continue
|
||||||
|
area: list[int] = cond_obj.area
|
||||||
|
if area is not None:
|
||||||
|
curr_default_mult: torch.Tensor = default_mults[i]
|
||||||
|
dims = len(area) // 2
|
||||||
|
for i in range(dims):
|
||||||
|
curr_default_mult = curr_default_mult.narrow(i + 2, area[i + dims], area[i])
|
||||||
|
curr_default_mult -= cond_obj.mult
|
||||||
|
else:
|
||||||
|
default_mults[i] -= cond_obj.mult
|
||||||
|
# for each default_mult, ReLU to make negatives=0, and then check for any nonzeros
|
||||||
|
for i, mult in enumerate(default_mults):
|
||||||
|
# if no default_cond for cond type, do nothing
|
||||||
|
if len(default_conds[i]) == 0:
|
||||||
|
continue
|
||||||
|
torch.nn.functional.relu(mult, inplace=True)
|
||||||
|
# if mult is all zeros, then don't add default_cond
|
||||||
|
if torch.max(mult) == 0.0:
|
||||||
|
continue
|
||||||
|
|
||||||
|
cond = default_conds[i]
|
||||||
|
for x in cond:
|
||||||
|
# do get_area_and_mult to get all the expected values
|
||||||
|
p = comfy.samplers.get_area_and_mult(x, x_in, timestep)
|
||||||
|
if p is None:
|
||||||
|
continue
|
||||||
|
# replace p's mult with calculated mult
|
||||||
|
p = p._replace(mult=mult)
|
||||||
|
if p.hooks is not None:
|
||||||
|
model.current_patcher.prepare_hook_patches_current_keyframe(timestep, p.hooks)
|
||||||
|
hooked_to_run.setdefault(p.hooks, list())
|
||||||
|
hooked_to_run[p.hooks] += [(p, i)]
|
||||||
|
|
||||||
|
def calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: torch.Tensor, timestep, model_options):
|
||||||
|
executor = comfy.patcher_extension.WrapperExecutor.new_executor(
|
||||||
|
_calc_cond_batch,
|
||||||
|
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.CALC_COND_BATCH, model_options, is_model_options=True)
|
||||||
|
)
|
||||||
|
return executor.execute(model, conds, x_in, timestep, model_options)
|
||||||
|
|
||||||
|
def _calc_cond_batch(model: 'BaseModel', conds: list[list[dict]], x_in: torch.Tensor, timestep, model_options):
|
||||||
out_conds = []
|
out_conds = []
|
||||||
out_counts = []
|
out_counts = []
|
||||||
to_run = []
|
# separate conds by matching hooks
|
||||||
|
hooked_to_run: dict[comfy.hooks.HookGroup,list[tuple[tuple,int]]] = {}
|
||||||
|
default_conds = []
|
||||||
|
has_default_conds = False
|
||||||
|
|
||||||
for i in range(len(conds)):
|
for i in range(len(conds)):
|
||||||
out_conds.append(torch.zeros_like(x_in))
|
out_conds.append(torch.zeros_like(x_in))
|
||||||
out_counts.append(torch.ones_like(x_in) * 1e-37)
|
out_counts.append(torch.ones_like(x_in) * 1e-37)
|
||||||
|
|
||||||
cond = conds[i]
|
cond = conds[i]
|
||||||
|
default_c = []
|
||||||
if cond is not None:
|
if cond is not None:
|
||||||
for x in cond:
|
for x in cond:
|
||||||
p = get_area_and_mult(x, x_in, timestep)
|
if 'default' in x:
|
||||||
|
default_c.append(x)
|
||||||
|
has_default_conds = True
|
||||||
|
continue
|
||||||
|
p = comfy.samplers.get_area_and_mult(x, x_in, timestep)
|
||||||
if p is None:
|
if p is None:
|
||||||
continue
|
continue
|
||||||
|
if p.hooks is not None:
|
||||||
|
model.current_patcher.prepare_hook_patches_current_keyframe(timestep, p.hooks)
|
||||||
|
hooked_to_run.setdefault(p.hooks, list())
|
||||||
|
hooked_to_run[p.hooks] += [(p, i)]
|
||||||
|
default_conds.append(default_c)
|
||||||
|
|
||||||
to_run += [(p, i)]
|
if has_default_conds:
|
||||||
|
finalize_default_conds(model, hooked_to_run, default_conds, x_in, timestep)
|
||||||
|
|
||||||
while len(to_run) > 0:
|
model.current_patcher.prepare_state(timestep)
|
||||||
first = to_run[0]
|
|
||||||
first_shape = first[0][0].shape
|
|
||||||
to_batch_temp = []
|
|
||||||
for x in range(len(to_run)):
|
|
||||||
if can_concat_cond(to_run[x][0], first[0]):
|
|
||||||
to_batch_temp += [x]
|
|
||||||
|
|
||||||
to_batch_temp.reverse()
|
# run every hooked_to_run separately
|
||||||
to_batch = to_batch_temp[:1]
|
for hooks, to_run in hooked_to_run.items():
|
||||||
|
while len(to_run) > 0:
|
||||||
|
first = to_run[0]
|
||||||
|
first_shape = first[0][0].shape
|
||||||
|
to_batch_temp = []
|
||||||
|
for x in range(len(to_run)):
|
||||||
|
if can_concat_cond(to_run[x][0], first[0]):
|
||||||
|
to_batch_temp += [x]
|
||||||
|
|
||||||
free_memory = model_management.get_free_memory(x_in.device)
|
to_batch_temp.reverse()
|
||||||
for i in range(1, len(to_batch_temp) + 1):
|
to_batch = to_batch_temp[:1]
|
||||||
batch_amount = to_batch_temp[:len(to_batch_temp)//i]
|
|
||||||
input_shape = [len(batch_amount) * first_shape[0]] + list(first_shape)[1:]
|
|
||||||
if model.memory_required(input_shape) * 1.5 < free_memory:
|
|
||||||
to_batch = batch_amount
|
|
||||||
break
|
|
||||||
|
|
||||||
input_x = []
|
free_memory = model_management.get_free_memory(x_in.device)
|
||||||
mult = []
|
for i in range(1, len(to_batch_temp) + 1):
|
||||||
c = []
|
batch_amount = to_batch_temp[:len(to_batch_temp)//i]
|
||||||
cond_or_uncond = []
|
input_shape = [len(batch_amount) * first_shape[0]] + list(first_shape)[1:]
|
||||||
area = []
|
if model.memory_required(input_shape) * 1.5 < free_memory:
|
||||||
control = None
|
to_batch = batch_amount
|
||||||
patches = None
|
break
|
||||||
for x in to_batch:
|
|
||||||
o = to_run.pop(x)
|
|
||||||
p = o[0]
|
|
||||||
input_x.append(p.input_x)
|
|
||||||
mult.append(p.mult)
|
|
||||||
c.append(p.conditioning)
|
|
||||||
area.append(p.area)
|
|
||||||
cond_or_uncond.append(o[1])
|
|
||||||
control = p.control
|
|
||||||
patches = p.patches
|
|
||||||
|
|
||||||
batch_chunks = len(cond_or_uncond)
|
input_x = []
|
||||||
input_x = torch.cat(input_x)
|
mult = []
|
||||||
c = cond_cat(c)
|
c = []
|
||||||
timestep_ = torch.cat([timestep] * batch_chunks)
|
cond_or_uncond = []
|
||||||
|
uuids = []
|
||||||
|
area = []
|
||||||
|
control = None
|
||||||
|
patches = None
|
||||||
|
for x in to_batch:
|
||||||
|
o = to_run.pop(x)
|
||||||
|
p = o[0]
|
||||||
|
input_x.append(p.input_x)
|
||||||
|
mult.append(p.mult)
|
||||||
|
c.append(p.conditioning)
|
||||||
|
area.append(p.area)
|
||||||
|
cond_or_uncond.append(o[1])
|
||||||
|
uuids.append(p.uuid)
|
||||||
|
control = p.control
|
||||||
|
patches = p.patches
|
||||||
|
|
||||||
if control is not None:
|
batch_chunks = len(cond_or_uncond)
|
||||||
c['control'] = control.get_control(input_x, timestep_, c, len(cond_or_uncond))
|
input_x = torch.cat(input_x)
|
||||||
|
c = cond_cat(c)
|
||||||
|
timestep_ = torch.cat([timestep] * batch_chunks)
|
||||||
|
|
||||||
transformer_options = {}
|
transformer_options = model.current_patcher.apply_hooks(hooks=hooks)
|
||||||
if 'transformer_options' in model_options:
|
if 'transformer_options' in model_options:
|
||||||
transformer_options = model_options['transformer_options'].copy()
|
transformer_options = comfy.patcher_extension.merge_nested_dicts(transformer_options,
|
||||||
|
model_options['transformer_options'],
|
||||||
|
copy_dict1=False)
|
||||||
|
|
||||||
if patches is not None:
|
if patches is not None:
|
||||||
if "patches" in transformer_options:
|
# TODO: replace with merge_nested_dicts function
|
||||||
cur_patches = transformer_options["patches"].copy()
|
if "patches" in transformer_options:
|
||||||
for p in patches:
|
cur_patches = transformer_options["patches"].copy()
|
||||||
if p in cur_patches:
|
for p in patches:
|
||||||
cur_patches[p] = cur_patches[p] + patches[p]
|
if p in cur_patches:
|
||||||
else:
|
cur_patches[p] = cur_patches[p] + patches[p]
|
||||||
cur_patches[p] = patches[p]
|
else:
|
||||||
transformer_options["patches"] = cur_patches
|
cur_patches[p] = patches[p]
|
||||||
|
transformer_options["patches"] = cur_patches
|
||||||
|
else:
|
||||||
|
transformer_options["patches"] = patches
|
||||||
|
|
||||||
|
transformer_options["cond_or_uncond"] = cond_or_uncond[:]
|
||||||
|
transformer_options["uuids"] = uuids[:]
|
||||||
|
transformer_options["sigmas"] = timestep
|
||||||
|
|
||||||
|
c['transformer_options'] = transformer_options
|
||||||
|
|
||||||
|
if control is not None:
|
||||||
|
c['control'] = control.get_control(input_x, timestep_, c, len(cond_or_uncond), transformer_options)
|
||||||
|
|
||||||
|
if 'model_function_wrapper' in model_options:
|
||||||
|
output = model_options['model_function_wrapper'](model.apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}).chunk(batch_chunks)
|
||||||
else:
|
else:
|
||||||
transformer_options["patches"] = patches
|
output = model.apply_model(input_x, timestep_, **c).chunk(batch_chunks)
|
||||||
|
|
||||||
transformer_options["cond_or_uncond"] = cond_or_uncond[:]
|
for o in range(batch_chunks):
|
||||||
transformer_options["sigmas"] = timestep
|
cond_index = cond_or_uncond[o]
|
||||||
|
a = area[o]
|
||||||
c['transformer_options'] = transformer_options
|
if a is None:
|
||||||
|
out_conds[cond_index] += output[o] * mult[o]
|
||||||
if 'model_function_wrapper' in model_options:
|
out_counts[cond_index] += mult[o]
|
||||||
output = model_options['model_function_wrapper'](model.apply_model, {"input": input_x, "timestep": timestep_, "c": c, "cond_or_uncond": cond_or_uncond}).chunk(batch_chunks)
|
else:
|
||||||
else:
|
out_c = out_conds[cond_index]
|
||||||
output = model.apply_model(input_x, timestep_, **c).chunk(batch_chunks)
|
out_cts = out_counts[cond_index]
|
||||||
|
dims = len(a) // 2
|
||||||
for o in range(batch_chunks):
|
for i in range(dims):
|
||||||
cond_index = cond_or_uncond[o]
|
out_c = out_c.narrow(i + 2, a[i + dims], a[i])
|
||||||
a = area[o]
|
out_cts = out_cts.narrow(i + 2, a[i + dims], a[i])
|
||||||
if a is None:
|
out_c += output[o] * mult[o]
|
||||||
out_conds[cond_index] += output[o] * mult[o]
|
out_cts += mult[o]
|
||||||
out_counts[cond_index] += mult[o]
|
|
||||||
else:
|
|
||||||
out_c = out_conds[cond_index]
|
|
||||||
out_cts = out_counts[cond_index]
|
|
||||||
dims = len(a) // 2
|
|
||||||
for i in range(dims):
|
|
||||||
out_c = out_c.narrow(i + 2, a[i + dims], a[i])
|
|
||||||
out_cts = out_cts.narrow(i + 2, a[i + dims], a[i])
|
|
||||||
out_c += output[o] * mult[o]
|
|
||||||
out_cts += mult[o]
|
|
||||||
|
|
||||||
for i in range(len(out_conds)):
|
for i in range(len(out_conds)):
|
||||||
out_conds[i] /= out_counts[i]
|
out_conds[i] /= out_counts[i]
|
||||||
@ -261,7 +346,7 @@ def cfg_function(model, cond_pred, uncond_pred, cond_scale, x, timestep, model_o
|
|||||||
cfg_result = uncond_pred + (cond_pred - uncond_pred) * cond_scale
|
cfg_result = uncond_pred + (cond_pred - uncond_pred) * cond_scale
|
||||||
|
|
||||||
for fn in model_options.get("sampler_post_cfg_function", []):
|
for fn in model_options.get("sampler_post_cfg_function", []):
|
||||||
args = {"denoised": cfg_result, "cond": cond, "uncond": uncond, "model": model, "uncond_denoised": uncond_pred, "cond_denoised": cond_pred,
|
args = {"denoised": cfg_result, "cond": cond, "uncond": uncond, "cond_scale": cond_scale, "model": model, "uncond_denoised": uncond_pred, "cond_denoised": cond_pred,
|
||||||
"sigma": timestep, "model_options": model_options, "input": x}
|
"sigma": timestep, "model_options": model_options, "input": x}
|
||||||
cfg_result = fn(args)
|
cfg_result = fn(args)
|
||||||
|
|
||||||
@ -500,10 +585,15 @@ def calculate_start_end_timesteps(model, conds):
|
|||||||
|
|
||||||
timestep_start = None
|
timestep_start = None
|
||||||
timestep_end = None
|
timestep_end = None
|
||||||
if 'start_percent' in x:
|
# handle clip hook schedule, if needed
|
||||||
timestep_start = s.percent_to_sigma(x['start_percent'])
|
if 'clip_start_percent' in x:
|
||||||
if 'end_percent' in x:
|
timestep_start = s.percent_to_sigma(max(x['clip_start_percent'], x.get('start_percent', 0.0)))
|
||||||
timestep_end = s.percent_to_sigma(x['end_percent'])
|
timestep_end = s.percent_to_sigma(min(x['clip_end_percent'], x.get('end_percent', 1.0)))
|
||||||
|
else:
|
||||||
|
if 'start_percent' in x:
|
||||||
|
timestep_start = s.percent_to_sigma(x['start_percent'])
|
||||||
|
if 'end_percent' in x:
|
||||||
|
timestep_end = s.percent_to_sigma(x['end_percent'])
|
||||||
|
|
||||||
if (timestep_start is not None) or (timestep_end is not None):
|
if (timestep_start is not None) or (timestep_end is not None):
|
||||||
n = x.copy()
|
n = x.copy()
|
||||||
@ -673,6 +763,12 @@ def process_conds(model, noise, conds, device, latent_image=None, denoise_mask=N
|
|||||||
if k != kk:
|
if k != kk:
|
||||||
create_cond_with_same_area_if_none(conds[kk], c)
|
create_cond_with_same_area_if_none(conds[kk], c)
|
||||||
|
|
||||||
|
for k in conds:
|
||||||
|
for c in conds[k]:
|
||||||
|
if 'hooks' in c:
|
||||||
|
for hook in c['hooks'].hooks:
|
||||||
|
hook.initialize_timesteps(model)
|
||||||
|
|
||||||
for k in conds:
|
for k in conds:
|
||||||
pre_run_control(model, conds[k])
|
pre_run_control(model, conds[k])
|
||||||
|
|
||||||
@ -685,9 +781,46 @@ def process_conds(model, noise, conds, device, latent_image=None, denoise_mask=N
|
|||||||
|
|
||||||
return conds
|
return conds
|
||||||
|
|
||||||
|
|
||||||
|
def preprocess_conds_hooks(conds: dict[str, list[dict[str]]]):
|
||||||
|
# determine which ControlNets have extra_hooks that should be combined with normal hooks
|
||||||
|
hook_replacement: dict[tuple[ControlBase, comfy.hooks.HookGroup], list[dict]] = {}
|
||||||
|
for k in conds:
|
||||||
|
for kk in conds[k]:
|
||||||
|
if 'control' in kk:
|
||||||
|
control: 'ControlBase' = kk['control']
|
||||||
|
extra_hooks = control.get_extra_hooks()
|
||||||
|
if len(extra_hooks) > 0:
|
||||||
|
hooks: comfy.hooks.HookGroup = kk.get('hooks', None)
|
||||||
|
to_replace = hook_replacement.setdefault((control, hooks), [])
|
||||||
|
to_replace.append(kk)
|
||||||
|
# if nothing to replace, do nothing
|
||||||
|
if len(hook_replacement) == 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
# for optimal sampling performance, common ControlNets + hook combos should have identical hooks
|
||||||
|
# on the cond dicts
|
||||||
|
for key, conds_to_modify in hook_replacement.items():
|
||||||
|
control = key[0]
|
||||||
|
hooks = key[1]
|
||||||
|
hooks = comfy.hooks.HookGroup.combine_all_hooks(control.get_extra_hooks() + [hooks])
|
||||||
|
# if combined hooks are not None, set as new hooks for all relevant conds
|
||||||
|
if hooks is not None:
|
||||||
|
for cond in conds_to_modify:
|
||||||
|
cond['hooks'] = hooks
|
||||||
|
|
||||||
|
|
||||||
|
def get_total_hook_groups_in_conds(conds: dict[str, list[dict[str]]]):
|
||||||
|
hooks_set = set()
|
||||||
|
for k in conds:
|
||||||
|
for kk in conds[k]:
|
||||||
|
hooks_set.add(kk.get('hooks', None))
|
||||||
|
return len(hooks_set)
|
||||||
|
|
||||||
|
|
||||||
class CFGGuider:
|
class CFGGuider:
|
||||||
def __init__(self, model_patcher):
|
def __init__(self, model_patcher):
|
||||||
self.model_patcher = model_patcher
|
self.model_patcher: 'ModelPatcher' = model_patcher
|
||||||
self.model_options = model_patcher.model_options
|
self.model_options = model_patcher.model_options
|
||||||
self.original_conds = {}
|
self.original_conds = {}
|
||||||
self.cfg = 1.0
|
self.cfg = 1.0
|
||||||
@ -714,19 +847,17 @@ class CFGGuider:
|
|||||||
|
|
||||||
self.conds = process_conds(self.inner_model, noise, self.conds, device, latent_image, denoise_mask, seed)
|
self.conds = process_conds(self.inner_model, noise, self.conds, device, latent_image, denoise_mask, seed)
|
||||||
|
|
||||||
extra_args = {"model_options": self.model_options, "seed":seed}
|
extra_args = {"model_options": comfy.model_patcher.create_model_options_clone(self.model_options), "seed": seed}
|
||||||
|
|
||||||
samples = sampler.sample(self, sigmas, extra_args, callback, noise, latent_image, denoise_mask, disable_pbar)
|
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||||
|
sampler.sample,
|
||||||
|
sampler,
|
||||||
|
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.SAMPLER_SAMPLE, extra_args["model_options"], is_model_options=True)
|
||||||
|
)
|
||||||
|
samples = executor.execute(self, sigmas, extra_args, callback, noise, latent_image, denoise_mask, disable_pbar)
|
||||||
return self.inner_model.process_latent_out(samples.to(torch.float32))
|
return self.inner_model.process_latent_out(samples.to(torch.float32))
|
||||||
|
|
||||||
def sample(self, noise, latent_image, sampler, sigmas, denoise_mask=None, callback=None, disable_pbar=False, seed=None):
|
def outer_sample(self, noise, latent_image, sampler, sigmas, denoise_mask=None, callback=None, disable_pbar=False, seed=None):
|
||||||
if sigmas.shape[-1] == 0:
|
|
||||||
return latent_image
|
|
||||||
|
|
||||||
self.conds = {}
|
|
||||||
for k in self.original_conds:
|
|
||||||
self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k]))
|
|
||||||
|
|
||||||
self.inner_model, self.conds, self.loaded_models = comfy.sampler_helpers.prepare_sampling(self.model_patcher, noise.shape, self.conds)
|
self.inner_model, self.conds, self.loaded_models = comfy.sampler_helpers.prepare_sampling(self.model_patcher, noise.shape, self.conds)
|
||||||
device = self.model_patcher.load_device
|
device = self.model_patcher.load_device
|
||||||
|
|
||||||
@ -737,14 +868,48 @@ class CFGGuider:
|
|||||||
latent_image = latent_image.to(device)
|
latent_image = latent_image.to(device)
|
||||||
sigmas = sigmas.to(device)
|
sigmas = sigmas.to(device)
|
||||||
|
|
||||||
output = self.inner_sample(noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed)
|
try:
|
||||||
|
self.model_patcher.pre_run()
|
||||||
|
output = self.inner_sample(noise, latent_image, device, sampler, sigmas, denoise_mask, callback, disable_pbar, seed)
|
||||||
|
finally:
|
||||||
|
self.model_patcher.cleanup()
|
||||||
|
|
||||||
comfy.sampler_helpers.cleanup_models(self.conds, self.loaded_models)
|
comfy.sampler_helpers.cleanup_models(self.conds, self.loaded_models)
|
||||||
del self.inner_model
|
del self.inner_model
|
||||||
del self.conds
|
|
||||||
del self.loaded_models
|
del self.loaded_models
|
||||||
return output
|
return output
|
||||||
|
|
||||||
|
def sample(self, noise, latent_image, sampler, sigmas, denoise_mask=None, callback=None, disable_pbar=False, seed=None):
|
||||||
|
if sigmas.shape[-1] == 0:
|
||||||
|
return latent_image
|
||||||
|
|
||||||
|
self.conds = {}
|
||||||
|
for k in self.original_conds:
|
||||||
|
self.conds[k] = list(map(lambda a: a.copy(), self.original_conds[k]))
|
||||||
|
preprocess_conds_hooks(self.conds)
|
||||||
|
|
||||||
|
try:
|
||||||
|
orig_model_options = self.model_options
|
||||||
|
self.model_options = comfy.model_patcher.create_model_options_clone(self.model_options)
|
||||||
|
# if one hook type (or just None), then don't bother caching weights for hooks (will never change after first step)
|
||||||
|
orig_hook_mode = self.model_patcher.hook_mode
|
||||||
|
if get_total_hook_groups_in_conds(self.conds) <= 1:
|
||||||
|
self.model_patcher.hook_mode = comfy.hooks.EnumHookMode.MinVram
|
||||||
|
comfy.sampler_helpers.prepare_model_patcher(self.model_patcher, self.conds, self.model_options)
|
||||||
|
executor = comfy.patcher_extension.WrapperExecutor.new_class_executor(
|
||||||
|
self.outer_sample,
|
||||||
|
self,
|
||||||
|
comfy.patcher_extension.get_all_wrappers(comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, self.model_options, is_model_options=True)
|
||||||
|
)
|
||||||
|
output = executor.execute(noise, latent_image, sampler, sigmas, denoise_mask, callback, disable_pbar, seed)
|
||||||
|
finally:
|
||||||
|
self.model_options = orig_model_options
|
||||||
|
self.model_patcher.hook_mode = orig_hook_mode
|
||||||
|
self.model_patcher.restore_hook_patches()
|
||||||
|
|
||||||
|
del self.conds
|
||||||
|
return output
|
||||||
|
|
||||||
|
|
||||||
def sample(model, noise, positive, negative, cfg, device, sampler, sigmas, model_options={}, latent_image=None, denoise_mask=None, callback=None, disable_pbar=False, seed=None):
|
def sample(model, noise, positive, negative, cfg, device, sampler, sigmas, model_options={}, latent_image=None, denoise_mask=None, callback=None, disable_pbar=False, seed=None):
|
||||||
cfg_guider = CFGGuider(model)
|
cfg_guider = CFGGuider(model)
|
||||||
|
|||||||
73
comfy/sd.py
73
comfy/sd.py
@ -1,8 +1,10 @@
|
|||||||
|
from __future__ import annotations
|
||||||
import torch
|
import torch
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
|
from comfy.utils import ProgressBar
|
||||||
from .ldm.models.autoencoder import AutoencoderKL, AutoencodingEngine
|
from .ldm.models.autoencoder import AutoencoderKL, AutoencodingEngine
|
||||||
from .ldm.cascade.stage_a import StageA
|
from .ldm.cascade.stage_a import StageA
|
||||||
from .ldm.cascade.stage_c_coder import StageC_coder
|
from .ldm.cascade.stage_c_coder import StageC_coder
|
||||||
@ -33,6 +35,7 @@ import comfy.text_encoders.lt
|
|||||||
import comfy.model_patcher
|
import comfy.model_patcher
|
||||||
import comfy.lora
|
import comfy.lora
|
||||||
import comfy.lora_convert
|
import comfy.lora_convert
|
||||||
|
import comfy.hooks
|
||||||
import comfy.t2i_adapter.adapter
|
import comfy.t2i_adapter.adapter
|
||||||
import comfy.taesd.taesd
|
import comfy.taesd.taesd
|
||||||
|
|
||||||
@ -98,9 +101,13 @@ class CLIP:
|
|||||||
|
|
||||||
self.tokenizer = tokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
self.tokenizer = tokenizer(embedding_directory=embedding_directory, tokenizer_data=tokenizer_data)
|
||||||
self.patcher = comfy.model_patcher.ModelPatcher(self.cond_stage_model, load_device=load_device, offload_device=offload_device)
|
self.patcher = comfy.model_patcher.ModelPatcher(self.cond_stage_model, load_device=load_device, offload_device=offload_device)
|
||||||
|
self.patcher.hook_mode = comfy.hooks.EnumHookMode.MinVram
|
||||||
|
self.patcher.is_clip = True
|
||||||
|
self.apply_hooks_to_conds = None
|
||||||
if params['device'] == load_device:
|
if params['device'] == load_device:
|
||||||
model_management.load_models_gpu([self.patcher], force_full_load=True)
|
model_management.load_models_gpu([self.patcher], force_full_load=True)
|
||||||
self.layer_idx = None
|
self.layer_idx = None
|
||||||
|
self.use_clip_schedule = False
|
||||||
logging.debug("CLIP model load device: {}, offload device: {}, current: {}".format(load_device, offload_device, params['device']))
|
logging.debug("CLIP model load device: {}, offload device: {}, current: {}".format(load_device, offload_device, params['device']))
|
||||||
|
|
||||||
def clone(self):
|
def clone(self):
|
||||||
@ -109,6 +116,8 @@ class CLIP:
|
|||||||
n.cond_stage_model = self.cond_stage_model
|
n.cond_stage_model = self.cond_stage_model
|
||||||
n.tokenizer = self.tokenizer
|
n.tokenizer = self.tokenizer
|
||||||
n.layer_idx = self.layer_idx
|
n.layer_idx = self.layer_idx
|
||||||
|
n.use_clip_schedule = self.use_clip_schedule
|
||||||
|
n.apply_hooks_to_conds = self.apply_hooks_to_conds
|
||||||
return n
|
return n
|
||||||
|
|
||||||
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
def add_patches(self, patches, strength_patch=1.0, strength_model=1.0):
|
||||||
@ -120,6 +129,69 @@ class CLIP:
|
|||||||
def tokenize(self, text, return_word_ids=False):
|
def tokenize(self, text, return_word_ids=False):
|
||||||
return self.tokenizer.tokenize_with_weights(text, return_word_ids)
|
return self.tokenizer.tokenize_with_weights(text, return_word_ids)
|
||||||
|
|
||||||
|
def add_hooks_to_dict(self, pooled_dict: dict[str]):
|
||||||
|
if self.apply_hooks_to_conds:
|
||||||
|
pooled_dict["hooks"] = self.apply_hooks_to_conds
|
||||||
|
return pooled_dict
|
||||||
|
|
||||||
|
def encode_from_tokens_scheduled(self, tokens, unprojected=False, add_dict: dict[str]={}, show_pbar=True):
|
||||||
|
all_cond_pooled: list[tuple[torch.Tensor, dict[str]]] = []
|
||||||
|
all_hooks = self.patcher.forced_hooks
|
||||||
|
if all_hooks is None or not self.use_clip_schedule:
|
||||||
|
# if no hooks or shouldn't use clip schedule, do unscheduled encode_from_tokens and perform add_dict
|
||||||
|
return_pooled = "unprojected" if unprojected else True
|
||||||
|
pooled_dict = self.encode_from_tokens(tokens, return_pooled=return_pooled, return_dict=True)
|
||||||
|
cond = pooled_dict.pop("cond")
|
||||||
|
# add/update any keys with the provided add_dict
|
||||||
|
pooled_dict.update(add_dict)
|
||||||
|
all_cond_pooled.append([cond, pooled_dict])
|
||||||
|
else:
|
||||||
|
scheduled_keyframes = all_hooks.get_hooks_for_clip_schedule()
|
||||||
|
|
||||||
|
self.cond_stage_model.reset_clip_options()
|
||||||
|
if self.layer_idx is not None:
|
||||||
|
self.cond_stage_model.set_clip_options({"layer": self.layer_idx})
|
||||||
|
if unprojected:
|
||||||
|
self.cond_stage_model.set_clip_options({"projected_pooled": False})
|
||||||
|
|
||||||
|
self.load_model()
|
||||||
|
all_hooks.reset()
|
||||||
|
self.patcher.patch_hooks(None)
|
||||||
|
if show_pbar:
|
||||||
|
pbar = ProgressBar(len(scheduled_keyframes))
|
||||||
|
|
||||||
|
for scheduled_opts in scheduled_keyframes:
|
||||||
|
t_range = scheduled_opts[0]
|
||||||
|
# don't bother encoding any conds outside of start_percent and end_percent bounds
|
||||||
|
if "start_percent" in add_dict:
|
||||||
|
if t_range[1] < add_dict["start_percent"]:
|
||||||
|
continue
|
||||||
|
if "end_percent" in add_dict:
|
||||||
|
if t_range[0] > add_dict["end_percent"]:
|
||||||
|
continue
|
||||||
|
hooks_keyframes = scheduled_opts[1]
|
||||||
|
for hook, keyframe in hooks_keyframes:
|
||||||
|
hook.hook_keyframe._current_keyframe = keyframe
|
||||||
|
# apply appropriate hooks with values that match new hook_keyframe
|
||||||
|
self.patcher.patch_hooks(all_hooks)
|
||||||
|
# perform encoding as normal
|
||||||
|
o = self.cond_stage_model.encode_token_weights(tokens)
|
||||||
|
cond, pooled = o[:2]
|
||||||
|
pooled_dict = {"pooled_output": pooled}
|
||||||
|
# add clip_start_percent and clip_end_percent in pooled
|
||||||
|
pooled_dict["clip_start_percent"] = t_range[0]
|
||||||
|
pooled_dict["clip_end_percent"] = t_range[1]
|
||||||
|
# add/update any keys with the provided add_dict
|
||||||
|
pooled_dict.update(add_dict)
|
||||||
|
# add hooks stored on clip
|
||||||
|
self.add_hooks_to_dict(pooled_dict)
|
||||||
|
all_cond_pooled.append([cond, pooled_dict])
|
||||||
|
if show_pbar:
|
||||||
|
pbar.update(1)
|
||||||
|
model_management.throw_exception_if_processing_interrupted()
|
||||||
|
all_hooks.reset()
|
||||||
|
return all_cond_pooled
|
||||||
|
|
||||||
def encode_from_tokens(self, tokens, return_pooled=False, return_dict=False):
|
def encode_from_tokens(self, tokens, return_pooled=False, return_dict=False):
|
||||||
self.cond_stage_model.reset_clip_options()
|
self.cond_stage_model.reset_clip_options()
|
||||||
|
|
||||||
@ -137,6 +209,7 @@ class CLIP:
|
|||||||
if len(o) > 2:
|
if len(o) > 2:
|
||||||
for k in o[2]:
|
for k in o[2]:
|
||||||
out[k] = o[2][k]
|
out[k] = o[2][k]
|
||||||
|
self.add_hooks_to_dict(out)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
if return_pooled:
|
if return_pooled:
|
||||||
|
|||||||
@ -90,8 +90,11 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
if textmodel_json_config is None:
|
if textmodel_json_config is None:
|
||||||
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "sd1_clip_config.json")
|
textmodel_json_config = os.path.join(os.path.dirname(os.path.realpath(__file__)), "sd1_clip_config.json")
|
||||||
|
|
||||||
with open(textmodel_json_config) as f:
|
if isinstance(textmodel_json_config, dict):
|
||||||
config = json.load(f)
|
config = textmodel_json_config
|
||||||
|
else:
|
||||||
|
with open(textmodel_json_config) as f:
|
||||||
|
config = json.load(f)
|
||||||
|
|
||||||
operations = model_options.get("custom_operations", None)
|
operations = model_options.get("custom_operations", None)
|
||||||
scaled_fp8 = None
|
scaled_fp8 = None
|
||||||
@ -196,11 +199,18 @@ class SDClipModel(torch.nn.Module, ClipTokenWeightEncoder):
|
|||||||
attention_mask = None
|
attention_mask = None
|
||||||
if self.enable_attention_masks or self.zero_out_masked or self.return_attention_masks:
|
if self.enable_attention_masks or self.zero_out_masked or self.return_attention_masks:
|
||||||
attention_mask = torch.zeros_like(tokens)
|
attention_mask = torch.zeros_like(tokens)
|
||||||
end_token = self.special_tokens.get("end", -1)
|
end_token = self.special_tokens.get("end", None)
|
||||||
|
if end_token is None:
|
||||||
|
cmp_token = self.special_tokens.get("pad", -1)
|
||||||
|
else:
|
||||||
|
cmp_token = end_token
|
||||||
|
|
||||||
for x in range(attention_mask.shape[0]):
|
for x in range(attention_mask.shape[0]):
|
||||||
for y in range(attention_mask.shape[1]):
|
for y in range(attention_mask.shape[1]):
|
||||||
attention_mask[x, y] = 1
|
attention_mask[x, y] = 1
|
||||||
if tokens[x, y] == end_token:
|
if tokens[x, y] == cmp_token:
|
||||||
|
if end_token is None:
|
||||||
|
attention_mask[x, y] = 0
|
||||||
break
|
break
|
||||||
|
|
||||||
attention_mask_model = None
|
attention_mask_model = None
|
||||||
@ -411,22 +421,25 @@ def load_embed(embedding_name, embedding_directory, embedding_size, embed_key=No
|
|||||||
return embed_out
|
return embed_out
|
||||||
|
|
||||||
class SDTokenizer:
|
class SDTokenizer:
|
||||||
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', tokenizer_class=CLIPTokenizer, has_start_token=True, pad_to_max_length=True, min_length=None, pad_token=None, tokenizer_data={}):
|
def __init__(self, tokenizer_path=None, max_length=77, pad_with_end=True, embedding_directory=None, embedding_size=768, embedding_key='clip_l', tokenizer_class=CLIPTokenizer, has_start_token=True, has_end_token=True, pad_to_max_length=True, min_length=None, pad_token=None, tokenizer_data={}):
|
||||||
if tokenizer_path is None:
|
if tokenizer_path is None:
|
||||||
tokenizer_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "sd1_tokenizer")
|
tokenizer_path = os.path.join(os.path.dirname(os.path.realpath(__file__)), "sd1_tokenizer")
|
||||||
self.tokenizer = tokenizer_class.from_pretrained(tokenizer_path)
|
self.tokenizer = tokenizer_class.from_pretrained(tokenizer_path)
|
||||||
self.max_length = max_length
|
self.max_length = max_length
|
||||||
self.min_length = min_length
|
self.min_length = min_length
|
||||||
|
self.end_token = None
|
||||||
|
|
||||||
empty = self.tokenizer('')["input_ids"]
|
empty = self.tokenizer('')["input_ids"]
|
||||||
if has_start_token:
|
if has_start_token:
|
||||||
self.tokens_start = 1
|
self.tokens_start = 1
|
||||||
self.start_token = empty[0]
|
self.start_token = empty[0]
|
||||||
self.end_token = empty[1]
|
if has_end_token:
|
||||||
|
self.end_token = empty[1]
|
||||||
else:
|
else:
|
||||||
self.tokens_start = 0
|
self.tokens_start = 0
|
||||||
self.start_token = None
|
self.start_token = None
|
||||||
self.end_token = empty[0]
|
if has_end_token:
|
||||||
|
self.end_token = empty[0]
|
||||||
|
|
||||||
if pad_token is not None:
|
if pad_token is not None:
|
||||||
self.pad_token = pad_token
|
self.pad_token = pad_token
|
||||||
@ -451,13 +464,16 @@ class SDTokenizer:
|
|||||||
Takes a potential embedding name and tries to retrieve it.
|
Takes a potential embedding name and tries to retrieve it.
|
||||||
Returns a Tuple consisting of the embedding and any leftover string, embedding can be None.
|
Returns a Tuple consisting of the embedding and any leftover string, embedding can be None.
|
||||||
'''
|
'''
|
||||||
|
split_embed = embedding_name.split(' ')
|
||||||
|
embedding_name = split_embed[0]
|
||||||
|
leftover = ' '.join(split_embed[1:])
|
||||||
embed = load_embed(embedding_name, self.embedding_directory, self.embedding_size, self.embedding_key)
|
embed = load_embed(embedding_name, self.embedding_directory, self.embedding_size, self.embedding_key)
|
||||||
if embed is None:
|
if embed is None:
|
||||||
stripped = embedding_name.strip(',')
|
stripped = embedding_name.strip(',')
|
||||||
if len(stripped) < len(embedding_name):
|
if len(stripped) < len(embedding_name):
|
||||||
embed = load_embed(stripped, self.embedding_directory, self.embedding_size, self.embedding_key)
|
embed = load_embed(stripped, self.embedding_directory, self.embedding_size, self.embedding_key)
|
||||||
return (embed, embedding_name[len(stripped):])
|
return (embed, "{} {}".format(embedding_name[len(stripped):], leftover))
|
||||||
return (embed, "")
|
return (embed, leftover)
|
||||||
|
|
||||||
|
|
||||||
def tokenize_with_weights(self, text:str, return_word_ids=False):
|
def tokenize_with_weights(self, text:str, return_word_ids=False):
|
||||||
@ -474,7 +490,12 @@ class SDTokenizer:
|
|||||||
#tokenize words
|
#tokenize words
|
||||||
tokens = []
|
tokens = []
|
||||||
for weighted_segment, weight in parsed_weights:
|
for weighted_segment, weight in parsed_weights:
|
||||||
to_tokenize = unescape_important(weighted_segment).replace("\n", " ").split(' ')
|
to_tokenize = unescape_important(weighted_segment).replace("\n", " ")
|
||||||
|
split = to_tokenize.split(' {}'.format(self.embedding_identifier))
|
||||||
|
to_tokenize = [split[0]]
|
||||||
|
for i in range(1, len(split)):
|
||||||
|
to_tokenize.append("{}{}".format(self.embedding_identifier, split[i]))
|
||||||
|
|
||||||
to_tokenize = [x for x in to_tokenize if x != ""]
|
to_tokenize = [x for x in to_tokenize if x != ""]
|
||||||
for word in to_tokenize:
|
for word in to_tokenize:
|
||||||
#if we find an embedding, deal with the embedding
|
#if we find an embedding, deal with the embedding
|
||||||
@ -493,8 +514,11 @@ class SDTokenizer:
|
|||||||
word = leftover
|
word = leftover
|
||||||
else:
|
else:
|
||||||
continue
|
continue
|
||||||
|
end = 999999999999
|
||||||
|
if self.end_token is not None:
|
||||||
|
end = -1
|
||||||
#parse word
|
#parse word
|
||||||
tokens.append([(t, weight) for t in self.tokenizer(word)["input_ids"][self.tokens_start:-1]])
|
tokens.append([(t, weight) for t in self.tokenizer(word)["input_ids"][self.tokens_start:end]])
|
||||||
|
|
||||||
#reshape token array to CLIP input size
|
#reshape token array to CLIP input size
|
||||||
batched_tokens = []
|
batched_tokens = []
|
||||||
@ -505,18 +529,24 @@ class SDTokenizer:
|
|||||||
for i, t_group in enumerate(tokens):
|
for i, t_group in enumerate(tokens):
|
||||||
#determine if we're going to try and keep the tokens in a single batch
|
#determine if we're going to try and keep the tokens in a single batch
|
||||||
is_large = len(t_group) >= self.max_word_length
|
is_large = len(t_group) >= self.max_word_length
|
||||||
|
if self.end_token is not None:
|
||||||
|
has_end_token = 1
|
||||||
|
else:
|
||||||
|
has_end_token = 0
|
||||||
|
|
||||||
while len(t_group) > 0:
|
while len(t_group) > 0:
|
||||||
if len(t_group) + len(batch) > self.max_length - 1:
|
if len(t_group) + len(batch) > self.max_length - has_end_token:
|
||||||
remaining_length = self.max_length - len(batch) - 1
|
remaining_length = self.max_length - len(batch) - has_end_token
|
||||||
#break word in two and add end token
|
#break word in two and add end token
|
||||||
if is_large:
|
if is_large:
|
||||||
batch.extend([(t,w,i+1) for t,w in t_group[:remaining_length]])
|
batch.extend([(t,w,i+1) for t,w in t_group[:remaining_length]])
|
||||||
batch.append((self.end_token, 1.0, 0))
|
if self.end_token is not None:
|
||||||
|
batch.append((self.end_token, 1.0, 0))
|
||||||
t_group = t_group[remaining_length:]
|
t_group = t_group[remaining_length:]
|
||||||
#add end token and pad
|
#add end token and pad
|
||||||
else:
|
else:
|
||||||
batch.append((self.end_token, 1.0, 0))
|
if self.end_token is not None:
|
||||||
|
batch.append((self.end_token, 1.0, 0))
|
||||||
if self.pad_to_max_length:
|
if self.pad_to_max_length:
|
||||||
batch.extend([(self.pad_token, 1.0, 0)] * (remaining_length))
|
batch.extend([(self.pad_token, 1.0, 0)] * (remaining_length))
|
||||||
#start new batch
|
#start new batch
|
||||||
@ -529,7 +559,8 @@ class SDTokenizer:
|
|||||||
t_group = []
|
t_group = []
|
||||||
|
|
||||||
#fill last batch
|
#fill last batch
|
||||||
batch.append((self.end_token, 1.0, 0))
|
if self.end_token is not None:
|
||||||
|
batch.append((self.end_token, 1.0, 0))
|
||||||
if self.pad_to_max_length:
|
if self.pad_to_max_length:
|
||||||
batch.extend([(self.pad_token, 1.0, 0)] * (self.max_length - len(batch)))
|
batch.extend([(self.pad_token, 1.0, 0)] * (self.max_length - len(batch)))
|
||||||
if self.min_length is not None and len(batch) < self.min_length:
|
if self.min_length is not None and len(batch) < self.min_length:
|
||||||
|
|||||||
@ -659,6 +659,15 @@ class Flux(supported_models_base.BASE):
|
|||||||
t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref))
|
t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref))
|
||||||
return supported_models_base.ClipTarget(comfy.text_encoders.flux.FluxTokenizer, comfy.text_encoders.flux.flux_clip(**t5_detect))
|
return supported_models_base.ClipTarget(comfy.text_encoders.flux.FluxTokenizer, comfy.text_encoders.flux.flux_clip(**t5_detect))
|
||||||
|
|
||||||
|
class FluxInpaint(Flux):
|
||||||
|
unet_config = {
|
||||||
|
"image_model": "flux",
|
||||||
|
"guidance_embed": True,
|
||||||
|
"in_channels": 96,
|
||||||
|
}
|
||||||
|
|
||||||
|
supported_inference_dtypes = [torch.bfloat16, torch.float32]
|
||||||
|
|
||||||
class FluxSchnell(Flux):
|
class FluxSchnell(Flux):
|
||||||
unet_config = {
|
unet_config = {
|
||||||
"image_model": "flux",
|
"image_model": "flux",
|
||||||
@ -731,6 +740,6 @@ class LTXV(supported_models_base.BASE):
|
|||||||
t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref))
|
t5_detect = comfy.text_encoders.sd3_clip.t5_xxl_detect(state_dict, "{}t5xxl.transformer.".format(pref))
|
||||||
return supported_models_base.ClipTarget(comfy.text_encoders.lt.LTXVT5Tokenizer, comfy.text_encoders.lt.ltxv_te(**t5_detect))
|
return supported_models_base.ClipTarget(comfy.text_encoders.lt.LTXVT5Tokenizer, comfy.text_encoders.lt.ltxv_te(**t5_detect))
|
||||||
|
|
||||||
models = [Stable_Zero123, SD15_instructpix2pix, SD15, SD20, SD21UnclipL, SD21UnclipH, SDXL_instructpix2pix, SDXLRefiner, SDXL, SSD1B, KOALA_700M, KOALA_1B, Segmind_Vega, SD_X4Upscaler, Stable_Cascade_C, Stable_Cascade_B, SV3D_u, SV3D_p, SD3, StableAudio, AuraFlow, HunyuanDiT, HunyuanDiT1, Flux, FluxSchnell, GenmoMochi, LTXV]
|
models = [Stable_Zero123, SD15_instructpix2pix, SD15, SD20, SD21UnclipL, SD21UnclipH, SDXL_instructpix2pix, SDXLRefiner, SDXL, SSD1B, KOALA_700M, KOALA_1B, Segmind_Vega, SD_X4Upscaler, Stable_Cascade_C, Stable_Cascade_B, SV3D_u, SV3D_p, SD3, StableAudio, AuraFlow, HunyuanDiT, HunyuanDiT1, FluxInpaint, Flux, FluxSchnell, GenmoMochi, LTXV]
|
||||||
|
|
||||||
models += [SVD_img2vid]
|
models += [SVD_img2vid]
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import os
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
class SPieceTokenizer:
|
class SPieceTokenizer:
|
||||||
|
|||||||
@ -209,6 +209,11 @@ class T5Stack(torch.nn.Module):
|
|||||||
intermediate = None
|
intermediate = None
|
||||||
optimized_attention = optimized_attention_for_device(x.device, mask=attention_mask is not None, small_input=True)
|
optimized_attention = optimized_attention_for_device(x.device, mask=attention_mask is not None, small_input=True)
|
||||||
past_bias = None
|
past_bias = None
|
||||||
|
|
||||||
|
if intermediate_output is not None:
|
||||||
|
if intermediate_output < 0:
|
||||||
|
intermediate_output = len(self.block) + intermediate_output
|
||||||
|
|
||||||
for i, l in enumerate(self.block):
|
for i, l in enumerate(self.block):
|
||||||
x, past_bias = l(x, mask, past_bias, optimized_attention)
|
x, past_bias = l(x, mask, past_bias, optimized_attention)
|
||||||
if i == intermediate_output:
|
if i == intermediate_output:
|
||||||
|
|||||||
@ -46,7 +46,13 @@ def load_torch_file(ckpt, safe_load=False, device=None):
|
|||||||
if "state_dict" in pl_sd:
|
if "state_dict" in pl_sd:
|
||||||
sd = pl_sd["state_dict"]
|
sd = pl_sd["state_dict"]
|
||||||
else:
|
else:
|
||||||
sd = pl_sd
|
if len(pl_sd) == 1:
|
||||||
|
key = list(pl_sd.keys())[0]
|
||||||
|
sd = pl_sd[key]
|
||||||
|
if not isinstance(sd, dict):
|
||||||
|
sd = pl_sd
|
||||||
|
else:
|
||||||
|
sd = pl_sd
|
||||||
return sd
|
return sd
|
||||||
|
|
||||||
def save_torch_file(sd, ckpt, metadata=None):
|
def save_torch_file(sd, ckpt, metadata=None):
|
||||||
|
|||||||
39
comfy_execution/validation.py
Normal file
39
comfy_execution/validation.py
Normal file
@ -0,0 +1,39 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
|
||||||
|
def validate_node_input(
|
||||||
|
received_type: str, input_type: str, strict: bool = False
|
||||||
|
) -> bool:
|
||||||
|
"""
|
||||||
|
received_type and input_type are both strings of the form "T1,T2,...".
|
||||||
|
|
||||||
|
If strict is True, the input_type must contain the received_type.
|
||||||
|
For example, if received_type is "STRING" and input_type is "STRING,INT",
|
||||||
|
this will return True. But if received_type is "STRING,INT" and input_type is
|
||||||
|
"INT", this will return False.
|
||||||
|
|
||||||
|
If strict is False, the input_type must have overlap with the received_type.
|
||||||
|
For example, if received_type is "STRING,BOOLEAN" and input_type is "STRING,INT",
|
||||||
|
this will return True.
|
||||||
|
|
||||||
|
Supports pre-union type extension behaviour of ``__ne__`` overrides.
|
||||||
|
"""
|
||||||
|
# If the types are exactly the same, we can return immediately
|
||||||
|
# Use pre-union behaviour: inverse of `__ne__`
|
||||||
|
if not received_type != input_type:
|
||||||
|
return True
|
||||||
|
|
||||||
|
# Not equal, and not strings
|
||||||
|
if not isinstance(received_type, str) or not isinstance(input_type, str):
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Split the type strings into sets for comparison
|
||||||
|
received_types = set(t.strip() for t in received_type.split(","))
|
||||||
|
input_types = set(t.strip() for t in input_type.split(","))
|
||||||
|
|
||||||
|
if strict:
|
||||||
|
# In strict mode, all received types must be in the input types
|
||||||
|
return received_types.issubset(input_types)
|
||||||
|
else:
|
||||||
|
# In non-strict mode, there must be at least one type in common
|
||||||
|
return len(received_types.intersection(input_types)) > 0
|
||||||
@ -2,8 +2,7 @@ import comfy.samplers
|
|||||||
import comfy.utils
|
import comfy.utils
|
||||||
import torch
|
import torch
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from tqdm.auto import trange, tqdm
|
from tqdm.auto import trange
|
||||||
import math
|
|
||||||
|
|
||||||
|
|
||||||
@torch.no_grad()
|
@torch.no_grad()
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import torch
|
|
||||||
from nodes import MAX_RESOLUTION
|
from nodes import MAX_RESOLUTION
|
||||||
|
|
||||||
class CLIPTextEncodeSDXLRefiner:
|
class CLIPTextEncodeSDXLRefiner:
|
||||||
@ -17,8 +16,7 @@ class CLIPTextEncodeSDXLRefiner:
|
|||||||
|
|
||||||
def encode(self, clip, ascore, width, height, text):
|
def encode(self, clip, ascore, width, height, text):
|
||||||
tokens = clip.tokenize(text)
|
tokens = clip.tokenize(text)
|
||||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
return (clip.encode_from_tokens_scheduled(tokens, add_dict={"aesthetic_score": ascore, "width": width, "height": height}), )
|
||||||
return ([[cond, {"pooled_output": pooled, "aesthetic_score": ascore, "width": width,"height": height}]], )
|
|
||||||
|
|
||||||
class CLIPTextEncodeSDXL:
|
class CLIPTextEncodeSDXL:
|
||||||
@classmethod
|
@classmethod
|
||||||
@ -47,8 +45,7 @@ class CLIPTextEncodeSDXL:
|
|||||||
tokens["l"] += empty["l"]
|
tokens["l"] += empty["l"]
|
||||||
while len(tokens["l"]) > len(tokens["g"]):
|
while len(tokens["l"]) > len(tokens["g"]):
|
||||||
tokens["g"] += empty["g"]
|
tokens["g"] += empty["g"]
|
||||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
return (clip.encode_from_tokens_scheduled(tokens, add_dict={"width": width, "height": height, "crop_w": crop_w, "crop_h": crop_h, "target_width": target_width, "target_height": target_height}), )
|
||||||
return ([[cond, {"pooled_output": pooled, "width": width, "height": height, "crop_w": crop_w, "crop_h": crop_h, "target_width": target_width, "target_height": target_height}]], )
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"CLIPTextEncodeSDXLRefiner": CLIPTextEncodeSDXLRefiner,
|
"CLIPTextEncodeSDXLRefiner": CLIPTextEncodeSDXLRefiner,
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import numpy as np
|
|
||||||
import torch
|
import torch
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
from enum import Enum
|
from enum import Enum
|
||||||
|
|||||||
@ -18,10 +18,7 @@ class CLIPTextEncodeFlux:
|
|||||||
tokens = clip.tokenize(clip_l)
|
tokens = clip.tokenize(clip_l)
|
||||||
tokens["t5xxl"] = clip.tokenize(t5xxl)["t5xxl"]
|
tokens["t5xxl"] = clip.tokenize(t5xxl)["t5xxl"]
|
||||||
|
|
||||||
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
|
return (clip.encode_from_tokens_scheduled(tokens, add_dict={"guidance": guidance}), )
|
||||||
cond = output.pop("cond")
|
|
||||||
output["guidance"] = guidance
|
|
||||||
return ([[cond, output]], )
|
|
||||||
|
|
||||||
class FluxGuidance:
|
class FluxGuidance:
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
744
comfy_extras/nodes_hooks.py
Normal file
744
comfy_extras/nodes_hooks.py
Normal file
@ -0,0 +1,744 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
from typing import TYPE_CHECKING, Union
|
||||||
|
import torch
|
||||||
|
from collections.abc import Iterable
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from comfy.sd import CLIP
|
||||||
|
|
||||||
|
import comfy.hooks
|
||||||
|
import comfy.sd
|
||||||
|
import comfy.utils
|
||||||
|
import folder_paths
|
||||||
|
|
||||||
|
###########################################
|
||||||
|
# Mask, Combine, and Hook Conditioning
|
||||||
|
#------------------------------------------
|
||||||
|
class PairConditioningSetProperties:
|
||||||
|
NodeId = 'PairConditioningSetProperties'
|
||||||
|
NodeName = 'Cond Pair Set Props'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"positive_NEW": ("CONDITIONING", ),
|
||||||
|
"negative_NEW": ("CONDITIONING", ),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"set_cond_area": (["default", "mask bounds"],),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"mask": ("MASK", ),
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
"timesteps": ("TIMESTEPS_RANGE",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||||
|
RETURN_NAMES = ("positive", "negative")
|
||||||
|
CATEGORY = "advanced/hooks/cond pair"
|
||||||
|
FUNCTION = "set_properties"
|
||||||
|
|
||||||
|
def set_properties(self, positive_NEW, negative_NEW,
|
||||||
|
strength: float, set_cond_area: str,
|
||||||
|
mask: torch.Tensor=None, hooks: comfy.hooks.HookGroup=None, timesteps: tuple=None):
|
||||||
|
final_positive, final_negative = comfy.hooks.set_conds_props(conds=[positive_NEW, negative_NEW],
|
||||||
|
strength=strength, set_cond_area=set_cond_area,
|
||||||
|
mask=mask, hooks=hooks, timesteps_range=timesteps)
|
||||||
|
return (final_positive, final_negative)
|
||||||
|
|
||||||
|
class PairConditioningSetPropertiesAndCombine:
|
||||||
|
NodeId = 'PairConditioningSetPropertiesAndCombine'
|
||||||
|
NodeName = 'Cond Pair Set Props Combine'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"positive": ("CONDITIONING", ),
|
||||||
|
"negative": ("CONDITIONING", ),
|
||||||
|
"positive_NEW": ("CONDITIONING", ),
|
||||||
|
"negative_NEW": ("CONDITIONING", ),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"set_cond_area": (["default", "mask bounds"],),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"mask": ("MASK", ),
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
"timesteps": ("TIMESTEPS_RANGE",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||||
|
RETURN_NAMES = ("positive", "negative")
|
||||||
|
CATEGORY = "advanced/hooks/cond pair"
|
||||||
|
FUNCTION = "set_properties"
|
||||||
|
|
||||||
|
def set_properties(self, positive, negative, positive_NEW, negative_NEW,
|
||||||
|
strength: float, set_cond_area: str,
|
||||||
|
mask: torch.Tensor=None, hooks: comfy.hooks.HookGroup=None, timesteps: tuple=None):
|
||||||
|
final_positive, final_negative = comfy.hooks.set_conds_props_and_combine(conds=[positive, negative], new_conds=[positive_NEW, negative_NEW],
|
||||||
|
strength=strength, set_cond_area=set_cond_area,
|
||||||
|
mask=mask, hooks=hooks, timesteps_range=timesteps)
|
||||||
|
return (final_positive, final_negative)
|
||||||
|
|
||||||
|
class ConditioningSetProperties:
|
||||||
|
NodeId = 'ConditioningSetProperties'
|
||||||
|
NodeName = 'Cond Set Props'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"cond_NEW": ("CONDITIONING", ),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"set_cond_area": (["default", "mask bounds"],),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"mask": ("MASK", ),
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
"timesteps": ("TIMESTEPS_RANGE",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
|
CATEGORY = "advanced/hooks/cond single"
|
||||||
|
FUNCTION = "set_properties"
|
||||||
|
|
||||||
|
def set_properties(self, cond_NEW,
|
||||||
|
strength: float, set_cond_area: str,
|
||||||
|
mask: torch.Tensor=None, hooks: comfy.hooks.HookGroup=None, timesteps: tuple=None):
|
||||||
|
(final_cond,) = comfy.hooks.set_conds_props(conds=[cond_NEW],
|
||||||
|
strength=strength, set_cond_area=set_cond_area,
|
||||||
|
mask=mask, hooks=hooks, timesteps_range=timesteps)
|
||||||
|
return (final_cond,)
|
||||||
|
|
||||||
|
class ConditioningSetPropertiesAndCombine:
|
||||||
|
NodeId = 'ConditioningSetPropertiesAndCombine'
|
||||||
|
NodeName = 'Cond Set Props Combine'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"cond": ("CONDITIONING", ),
|
||||||
|
"cond_NEW": ("CONDITIONING", ),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
|
"set_cond_area": (["default", "mask bounds"],),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"mask": ("MASK", ),
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
"timesteps": ("TIMESTEPS_RANGE",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
|
CATEGORY = "advanced/hooks/cond single"
|
||||||
|
FUNCTION = "set_properties"
|
||||||
|
|
||||||
|
def set_properties(self, cond, cond_NEW,
|
||||||
|
strength: float, set_cond_area: str,
|
||||||
|
mask: torch.Tensor=None, hooks: comfy.hooks.HookGroup=None, timesteps: tuple=None):
|
||||||
|
(final_cond,) = comfy.hooks.set_conds_props_and_combine(conds=[cond], new_conds=[cond_NEW],
|
||||||
|
strength=strength, set_cond_area=set_cond_area,
|
||||||
|
mask=mask, hooks=hooks, timesteps_range=timesteps)
|
||||||
|
return (final_cond,)
|
||||||
|
|
||||||
|
class PairConditioningCombine:
|
||||||
|
NodeId = 'PairConditioningCombine'
|
||||||
|
NodeName = 'Cond Pair Combine'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"positive_A": ("CONDITIONING",),
|
||||||
|
"negative_A": ("CONDITIONING",),
|
||||||
|
"positive_B": ("CONDITIONING",),
|
||||||
|
"negative_B": ("CONDITIONING",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||||
|
RETURN_NAMES = ("positive", "negative")
|
||||||
|
CATEGORY = "advanced/hooks/cond pair"
|
||||||
|
FUNCTION = "combine"
|
||||||
|
|
||||||
|
def combine(self, positive_A, negative_A, positive_B, negative_B):
|
||||||
|
final_positive, final_negative = comfy.hooks.set_conds_props_and_combine(conds=[positive_A, negative_A], new_conds=[positive_B, negative_B],)
|
||||||
|
return (final_positive, final_negative,)
|
||||||
|
|
||||||
|
class PairConditioningSetDefaultAndCombine:
|
||||||
|
NodeId = 'PairConditioningSetDefaultCombine'
|
||||||
|
NodeName = 'Cond Pair Set Default Combine'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"positive": ("CONDITIONING",),
|
||||||
|
"negative": ("CONDITIONING",),
|
||||||
|
"positive_DEFAULT": ("CONDITIONING",),
|
||||||
|
"negative_DEFAULT": ("CONDITIONING",),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING", "CONDITIONING")
|
||||||
|
RETURN_NAMES = ("positive", "negative")
|
||||||
|
CATEGORY = "advanced/hooks/cond pair"
|
||||||
|
FUNCTION = "set_default_and_combine"
|
||||||
|
|
||||||
|
def set_default_and_combine(self, positive, negative, positive_DEFAULT, negative_DEFAULT,
|
||||||
|
hooks: comfy.hooks.HookGroup=None):
|
||||||
|
final_positive, final_negative = comfy.hooks.set_default_conds_and_combine(conds=[positive, negative], new_conds=[positive_DEFAULT, negative_DEFAULT],
|
||||||
|
hooks=hooks)
|
||||||
|
return (final_positive, final_negative)
|
||||||
|
|
||||||
|
class ConditioningSetDefaultAndCombine:
|
||||||
|
NodeId = 'ConditioningSetDefaultCombine'
|
||||||
|
NodeName = 'Cond Set Default Combine'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"cond": ("CONDITIONING",),
|
||||||
|
"cond_DEFAULT": ("CONDITIONING",),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
|
CATEGORY = "advanced/hooks/cond single"
|
||||||
|
FUNCTION = "set_default_and_combine"
|
||||||
|
|
||||||
|
def set_default_and_combine(self, cond, cond_DEFAULT,
|
||||||
|
hooks: comfy.hooks.HookGroup=None):
|
||||||
|
(final_conditioning,) = comfy.hooks.set_default_conds_and_combine(conds=[cond], new_conds=[cond_DEFAULT],
|
||||||
|
hooks=hooks)
|
||||||
|
return (final_conditioning,)
|
||||||
|
|
||||||
|
class SetClipHooks:
|
||||||
|
NodeId = 'SetClipHooks'
|
||||||
|
NodeName = 'Set CLIP Hooks'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"clip": ("CLIP",),
|
||||||
|
"apply_to_conds": ("BOOLEAN", {"default": True}),
|
||||||
|
"schedule_clip": ("BOOLEAN", {"default": False})
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hooks": ("HOOKS",)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CLIP",)
|
||||||
|
CATEGORY = "advanced/hooks/clip"
|
||||||
|
FUNCTION = "apply_hooks"
|
||||||
|
|
||||||
|
def apply_hooks(self, clip: 'CLIP', schedule_clip: bool, apply_to_conds: bool, hooks: comfy.hooks.HookGroup=None):
|
||||||
|
if hooks is not None:
|
||||||
|
clip = clip.clone()
|
||||||
|
if apply_to_conds:
|
||||||
|
clip.apply_hooks_to_conds = hooks
|
||||||
|
clip.patcher.forced_hooks = hooks.clone()
|
||||||
|
clip.use_clip_schedule = schedule_clip
|
||||||
|
if not clip.use_clip_schedule:
|
||||||
|
clip.patcher.forced_hooks.set_keyframes_on_hooks(None)
|
||||||
|
clip.patcher.register_all_hook_patches(hooks.get_dict_repr(), comfy.hooks.EnumWeightTarget.Clip)
|
||||||
|
return (clip,)
|
||||||
|
|
||||||
|
class ConditioningTimestepsRange:
|
||||||
|
NodeId = 'ConditioningTimestepsRange'
|
||||||
|
NodeName = 'Timesteps Range'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001})
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("TIMESTEPS_RANGE", "TIMESTEPS_RANGE", "TIMESTEPS_RANGE")
|
||||||
|
RETURN_NAMES = ("TIMESTEPS_RANGE", "BEFORE_RANGE", "AFTER_RANGE")
|
||||||
|
CATEGORY = "advanced/hooks"
|
||||||
|
FUNCTION = "create_range"
|
||||||
|
|
||||||
|
def create_range(self, start_percent: float, end_percent: float):
|
||||||
|
return ((start_percent, end_percent), (0.0, start_percent), (end_percent, 1.0))
|
||||||
|
#------------------------------------------
|
||||||
|
###########################################
|
||||||
|
|
||||||
|
|
||||||
|
###########################################
|
||||||
|
# Create Hooks
|
||||||
|
#------------------------------------------
|
||||||
|
class CreateHookLora:
|
||||||
|
NodeId = 'CreateHookLora'
|
||||||
|
NodeName = 'Create Hook LoRA'
|
||||||
|
def __init__(self):
|
||||||
|
self.loaded_lora = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"lora_name": (folder_paths.get_filename_list("loras"), ),
|
||||||
|
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
"strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hooks": ("HOOKS",)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/create"
|
||||||
|
FUNCTION = "create_hook"
|
||||||
|
|
||||||
|
def create_hook(self, lora_name: str, strength_model: float, strength_clip: float, prev_hooks: comfy.hooks.HookGroup=None):
|
||||||
|
if prev_hooks is None:
|
||||||
|
prev_hooks = comfy.hooks.HookGroup()
|
||||||
|
prev_hooks.clone()
|
||||||
|
|
||||||
|
if strength_model == 0 and strength_clip == 0:
|
||||||
|
return (prev_hooks,)
|
||||||
|
|
||||||
|
lora_path = folder_paths.get_full_path("loras", lora_name)
|
||||||
|
lora = None
|
||||||
|
if self.loaded_lora is not None:
|
||||||
|
if self.loaded_lora[0] == lora_path:
|
||||||
|
lora = self.loaded_lora[1]
|
||||||
|
else:
|
||||||
|
temp = self.loaded_lora
|
||||||
|
self.loaded_lora = None
|
||||||
|
del temp
|
||||||
|
|
||||||
|
if lora is None:
|
||||||
|
lora = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||||
|
self.loaded_lora = (lora_path, lora)
|
||||||
|
|
||||||
|
hooks = comfy.hooks.create_hook_lora(lora=lora, strength_model=strength_model, strength_clip=strength_clip)
|
||||||
|
return (prev_hooks.clone_and_combine(hooks),)
|
||||||
|
|
||||||
|
class CreateHookLoraModelOnly(CreateHookLora):
|
||||||
|
NodeId = 'CreateHookLoraModelOnly'
|
||||||
|
NodeName = 'Create Hook LoRA (MO)'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"lora_name": (folder_paths.get_filename_list("loras"), ),
|
||||||
|
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hooks": ("HOOKS",)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/create"
|
||||||
|
FUNCTION = "create_hook_model_only"
|
||||||
|
|
||||||
|
def create_hook_model_only(self, lora_name: str, strength_model: float, prev_hooks: comfy.hooks.HookGroup=None):
|
||||||
|
return self.create_hook(lora_name=lora_name, strength_model=strength_model, strength_clip=0, prev_hooks=prev_hooks)
|
||||||
|
|
||||||
|
class CreateHookModelAsLora:
|
||||||
|
NodeId = 'CreateHookModelAsLora'
|
||||||
|
NodeName = 'Create Hook Model as LoRA'
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
# when not None, will be in following format:
|
||||||
|
# (ckpt_path: str, weights_model: dict, weights_clip: dict)
|
||||||
|
self.loaded_weights = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
|
||||||
|
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
"strength_clip": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hooks": ("HOOKS",)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/create"
|
||||||
|
FUNCTION = "create_hook"
|
||||||
|
|
||||||
|
def create_hook(self, ckpt_name: str, strength_model: float, strength_clip: float,
|
||||||
|
prev_hooks: comfy.hooks.HookGroup=None):
|
||||||
|
if prev_hooks is None:
|
||||||
|
prev_hooks = comfy.hooks.HookGroup()
|
||||||
|
prev_hooks.clone()
|
||||||
|
|
||||||
|
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
|
||||||
|
weights_model = None
|
||||||
|
weights_clip = None
|
||||||
|
if self.loaded_weights is not None:
|
||||||
|
if self.loaded_weights[0] == ckpt_path:
|
||||||
|
weights_model = self.loaded_weights[1]
|
||||||
|
weights_clip = self.loaded_weights[2]
|
||||||
|
else:
|
||||||
|
temp = self.loaded_weights
|
||||||
|
self.loaded_weights = None
|
||||||
|
del temp
|
||||||
|
|
||||||
|
if weights_model is None:
|
||||||
|
out = comfy.sd.load_checkpoint_guess_config(ckpt_path, output_vae=True, output_clip=True, embedding_directory=folder_paths.get_folder_paths("embeddings"))
|
||||||
|
weights_model = comfy.hooks.get_patch_weights_from_model(out[0])
|
||||||
|
weights_clip = comfy.hooks.get_patch_weights_from_model(out[1].patcher if out[1] else out[1])
|
||||||
|
self.loaded_weights = (ckpt_path, weights_model, weights_clip)
|
||||||
|
|
||||||
|
hooks = comfy.hooks.create_hook_model_as_lora(weights_model=weights_model, weights_clip=weights_clip,
|
||||||
|
strength_model=strength_model, strength_clip=strength_clip)
|
||||||
|
return (prev_hooks.clone_and_combine(hooks),)
|
||||||
|
|
||||||
|
class CreateHookModelAsLoraModelOnly(CreateHookModelAsLora):
|
||||||
|
NodeId = 'CreateHookModelAsLoraModelOnly'
|
||||||
|
NodeName = 'Create Hook Model as LoRA (MO)'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"ckpt_name": (folder_paths.get_filename_list("checkpoints"), ),
|
||||||
|
"strength_model": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hooks": ("HOOKS",)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/create"
|
||||||
|
FUNCTION = "create_hook_model_only"
|
||||||
|
|
||||||
|
def create_hook_model_only(self, ckpt_name: str, strength_model: float,
|
||||||
|
prev_hooks: comfy.hooks.HookGroup=None):
|
||||||
|
return self.create_hook(ckpt_name=ckpt_name, strength_model=strength_model, strength_clip=0.0, prev_hooks=prev_hooks)
|
||||||
|
#------------------------------------------
|
||||||
|
###########################################
|
||||||
|
|
||||||
|
|
||||||
|
###########################################
|
||||||
|
# Schedule Hooks
|
||||||
|
#------------------------------------------
|
||||||
|
class SetHookKeyframes:
|
||||||
|
NodeId = 'SetHookKeyframes'
|
||||||
|
NodeName = 'Set Hook Keyframes'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "set_hook_keyframes"
|
||||||
|
|
||||||
|
def set_hook_keyframes(self, hooks: comfy.hooks.HookGroup, hook_kf: comfy.hooks.HookKeyframeGroup=None):
|
||||||
|
if hook_kf is not None:
|
||||||
|
hooks = hooks.clone()
|
||||||
|
hooks.set_keyframes_on_hooks(hook_kf=hook_kf)
|
||||||
|
return (hooks,)
|
||||||
|
|
||||||
|
class CreateHookKeyframe:
|
||||||
|
NodeId = 'CreateHookKeyframe'
|
||||||
|
NodeName = 'Create Hook Keyframe'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"strength_mult": ("FLOAT", {"default": 1.0, "min": -20.0, "max": 20.0, "step": 0.01}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOK_KEYFRAMES",)
|
||||||
|
RETURN_NAMES = ("HOOK_KF",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "create_hook_keyframe"
|
||||||
|
|
||||||
|
def create_hook_keyframe(self, strength_mult: float, start_percent: float, prev_hook_kf: comfy.hooks.HookKeyframeGroup=None):
|
||||||
|
if prev_hook_kf is None:
|
||||||
|
prev_hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||||
|
prev_hook_kf = prev_hook_kf.clone()
|
||||||
|
keyframe = comfy.hooks.HookKeyframe(strength=strength_mult, start_percent=start_percent)
|
||||||
|
prev_hook_kf.add(keyframe)
|
||||||
|
return (prev_hook_kf,)
|
||||||
|
|
||||||
|
class CreateHookKeyframesInterpolated:
|
||||||
|
NodeId = 'CreateHookKeyframesInterpolated'
|
||||||
|
NodeName = 'Create Hook Keyframes Interp.'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"strength_start": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||||
|
"strength_end": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}, ),
|
||||||
|
"interpolation": (comfy.hooks.InterpolationMethod._LIST, ),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
"keyframes_count": ("INT", {"default": 5, "min": 2, "max": 100, "step": 1}),
|
||||||
|
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOK_KEYFRAMES",)
|
||||||
|
RETURN_NAMES = ("HOOK_KF",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "create_hook_keyframes"
|
||||||
|
|
||||||
|
def create_hook_keyframes(self, strength_start: float, strength_end: float, interpolation: str,
|
||||||
|
start_percent: float, end_percent: float, keyframes_count: int,
|
||||||
|
print_keyframes=False, prev_hook_kf: comfy.hooks.HookKeyframeGroup=None):
|
||||||
|
if prev_hook_kf is None:
|
||||||
|
prev_hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||||
|
prev_hook_kf = prev_hook_kf.clone()
|
||||||
|
percents = comfy.hooks.InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=keyframes_count,
|
||||||
|
method=comfy.hooks.InterpolationMethod.LINEAR)
|
||||||
|
strengths = comfy.hooks.InterpolationMethod.get_weights(num_from=strength_start, num_to=strength_end, length=keyframes_count, method=interpolation)
|
||||||
|
|
||||||
|
is_first = True
|
||||||
|
for percent, strength in zip(percents, strengths):
|
||||||
|
guarantee_steps = 0
|
||||||
|
if is_first:
|
||||||
|
guarantee_steps = 1
|
||||||
|
is_first = False
|
||||||
|
prev_hook_kf.add(comfy.hooks.HookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps))
|
||||||
|
if print_keyframes:
|
||||||
|
print(f"Hook Keyframe - start_percent:{percent} = {strength}")
|
||||||
|
return (prev_hook_kf,)
|
||||||
|
|
||||||
|
class CreateHookKeyframesFromFloats:
|
||||||
|
NodeId = 'CreateHookKeyframesFromFloats'
|
||||||
|
NodeName = 'Create Hook Keyframes From Floats'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"floats_strength": ("FLOATS", {"default": -1, "min": -1, "step": 0.001, "forceInput": True}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
"print_keyframes": ("BOOLEAN", {"default": False}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"prev_hook_kf": ("HOOK_KEYFRAMES",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOK_KEYFRAMES",)
|
||||||
|
RETURN_NAMES = ("HOOK_KF",)
|
||||||
|
CATEGORY = "advanced/hooks/scheduling"
|
||||||
|
FUNCTION = "create_hook_keyframes"
|
||||||
|
|
||||||
|
def create_hook_keyframes(self, floats_strength: Union[float, list[float]],
|
||||||
|
start_percent: float, end_percent: float,
|
||||||
|
prev_hook_kf: comfy.hooks.HookKeyframeGroup=None, print_keyframes=False):
|
||||||
|
if prev_hook_kf is None:
|
||||||
|
prev_hook_kf = comfy.hooks.HookKeyframeGroup()
|
||||||
|
prev_hook_kf = prev_hook_kf.clone()
|
||||||
|
if type(floats_strength) in (float, int):
|
||||||
|
floats_strength = [float(floats_strength)]
|
||||||
|
elif isinstance(floats_strength, Iterable):
|
||||||
|
pass
|
||||||
|
else:
|
||||||
|
raise Exception(f"floats_strength must be either an iterable input or a float, but was{type(floats_strength).__repr__}.")
|
||||||
|
percents = comfy.hooks.InterpolationMethod.get_weights(num_from=start_percent, num_to=end_percent, length=len(floats_strength),
|
||||||
|
method=comfy.hooks.InterpolationMethod.LINEAR)
|
||||||
|
|
||||||
|
is_first = True
|
||||||
|
for percent, strength in zip(percents, floats_strength):
|
||||||
|
guarantee_steps = 0
|
||||||
|
if is_first:
|
||||||
|
guarantee_steps = 1
|
||||||
|
is_first = False
|
||||||
|
prev_hook_kf.add(comfy.hooks.HookKeyframe(strength=strength, start_percent=percent, guarantee_steps=guarantee_steps))
|
||||||
|
if print_keyframes:
|
||||||
|
print(f"Hook Keyframe - start_percent:{percent} = {strength}")
|
||||||
|
return (prev_hook_kf,)
|
||||||
|
#------------------------------------------
|
||||||
|
###########################################
|
||||||
|
|
||||||
|
|
||||||
|
class SetModelHooksOnCond:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
"conditioning": ("CONDITIONING",),
|
||||||
|
"hooks": ("HOOKS",),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
|
CATEGORY = "advanced/hooks/manual"
|
||||||
|
FUNCTION = "attach_hook"
|
||||||
|
|
||||||
|
def attach_hook(self, conditioning, hooks: comfy.hooks.HookGroup):
|
||||||
|
return (comfy.hooks.set_hooks_for_conditioning(conditioning, hooks),)
|
||||||
|
|
||||||
|
|
||||||
|
###########################################
|
||||||
|
# Combine Hooks
|
||||||
|
#------------------------------------------
|
||||||
|
class CombineHooks:
|
||||||
|
NodeId = 'CombineHooks2'
|
||||||
|
NodeName = 'Combine Hooks [2]'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hooks_A": ("HOOKS",),
|
||||||
|
"hooks_B": ("HOOKS",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/combine"
|
||||||
|
FUNCTION = "combine_hooks"
|
||||||
|
|
||||||
|
def combine_hooks(self,
|
||||||
|
hooks_A: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_B: comfy.hooks.HookGroup=None):
|
||||||
|
candidates = [hooks_A, hooks_B]
|
||||||
|
return (comfy.hooks.HookGroup.combine_all_hooks(candidates),)
|
||||||
|
|
||||||
|
class CombineHooksFour:
|
||||||
|
NodeId = 'CombineHooks4'
|
||||||
|
NodeName = 'Combine Hooks [4]'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hooks_A": ("HOOKS",),
|
||||||
|
"hooks_B": ("HOOKS",),
|
||||||
|
"hooks_C": ("HOOKS",),
|
||||||
|
"hooks_D": ("HOOKS",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/combine"
|
||||||
|
FUNCTION = "combine_hooks"
|
||||||
|
|
||||||
|
def combine_hooks(self,
|
||||||
|
hooks_A: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_B: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_C: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_D: comfy.hooks.HookGroup=None):
|
||||||
|
candidates = [hooks_A, hooks_B, hooks_C, hooks_D]
|
||||||
|
return (comfy.hooks.HookGroup.combine_all_hooks(candidates),)
|
||||||
|
|
||||||
|
class CombineHooksEight:
|
||||||
|
NodeId = 'CombineHooks8'
|
||||||
|
NodeName = 'Combine Hooks [8]'
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {
|
||||||
|
"required": {
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"hooks_A": ("HOOKS",),
|
||||||
|
"hooks_B": ("HOOKS",),
|
||||||
|
"hooks_C": ("HOOKS",),
|
||||||
|
"hooks_D": ("HOOKS",),
|
||||||
|
"hooks_E": ("HOOKS",),
|
||||||
|
"hooks_F": ("HOOKS",),
|
||||||
|
"hooks_G": ("HOOKS",),
|
||||||
|
"hooks_H": ("HOOKS",),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
EXPERIMENTAL = True
|
||||||
|
RETURN_TYPES = ("HOOKS",)
|
||||||
|
CATEGORY = "advanced/hooks/combine"
|
||||||
|
FUNCTION = "combine_hooks"
|
||||||
|
|
||||||
|
def combine_hooks(self,
|
||||||
|
hooks_A: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_B: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_C: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_D: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_E: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_F: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_G: comfy.hooks.HookGroup=None,
|
||||||
|
hooks_H: comfy.hooks.HookGroup=None):
|
||||||
|
candidates = [hooks_A, hooks_B, hooks_C, hooks_D, hooks_E, hooks_F, hooks_G, hooks_H]
|
||||||
|
return (comfy.hooks.HookGroup.combine_all_hooks(candidates),)
|
||||||
|
#------------------------------------------
|
||||||
|
###########################################
|
||||||
|
|
||||||
|
node_list = [
|
||||||
|
# Create
|
||||||
|
CreateHookLora,
|
||||||
|
CreateHookLoraModelOnly,
|
||||||
|
CreateHookModelAsLora,
|
||||||
|
CreateHookModelAsLoraModelOnly,
|
||||||
|
# Scheduling
|
||||||
|
SetHookKeyframes,
|
||||||
|
CreateHookKeyframe,
|
||||||
|
CreateHookKeyframesInterpolated,
|
||||||
|
CreateHookKeyframesFromFloats,
|
||||||
|
# Combine
|
||||||
|
CombineHooks,
|
||||||
|
CombineHooksFour,
|
||||||
|
CombineHooksEight,
|
||||||
|
# Attach
|
||||||
|
ConditioningSetProperties,
|
||||||
|
ConditioningSetPropertiesAndCombine,
|
||||||
|
PairConditioningSetProperties,
|
||||||
|
PairConditioningSetPropertiesAndCombine,
|
||||||
|
ConditioningSetDefaultAndCombine,
|
||||||
|
PairConditioningSetDefaultAndCombine,
|
||||||
|
PairConditioningCombine,
|
||||||
|
SetClipHooks,
|
||||||
|
# Other
|
||||||
|
ConditioningTimestepsRange,
|
||||||
|
]
|
||||||
|
NODE_CLASS_MAPPINGS = {}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||||
|
|
||||||
|
for node in node_list:
|
||||||
|
NODE_CLASS_MAPPINGS[node.NodeId] = node
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS[node.NodeId] = node.NodeName
|
||||||
@ -15,9 +15,7 @@ class CLIPTextEncodeHunyuanDiT:
|
|||||||
tokens = clip.tokenize(bert)
|
tokens = clip.tokenize(bert)
|
||||||
tokens["mt5xl"] = clip.tokenize(mt5xl)["mt5xl"]
|
tokens["mt5xl"] = clip.tokenize(mt5xl)["mt5xl"]
|
||||||
|
|
||||||
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
|
return (clip.encode_from_tokens_scheduled(tokens), )
|
||||||
cond = output.pop("cond")
|
|
||||||
return ([[cond, output]], )
|
|
||||||
|
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
|||||||
@ -32,7 +32,9 @@ class LTXVImgToVideo:
|
|||||||
"width": ("INT", {"default": 768, "min": 64, "max": nodes.MAX_RESOLUTION, "step": 32}),
|
"width": ("INT", {"default": 768, "min": 64, "max": nodes.MAX_RESOLUTION, "step": 32}),
|
||||||
"height": ("INT", {"default": 512, "min": 64, "max": nodes.MAX_RESOLUTION, "step": 32}),
|
"height": ("INT", {"default": 512, "min": 64, "max": nodes.MAX_RESOLUTION, "step": 32}),
|
||||||
"length": ("INT", {"default": 97, "min": 9, "max": nodes.MAX_RESOLUTION, "step": 8}),
|
"length": ("INT", {"default": 97, "min": 9, "max": nodes.MAX_RESOLUTION, "step": 8}),
|
||||||
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096})}}
|
"batch_size": ("INT", {"default": 1, "min": 1, "max": 4096}),
|
||||||
|
"image_noise_scale": ("FLOAT", {"default": 0.15, "min": 0, "max": 1.0, "step": 0.01, "tooltip": "Amount of noise to apply on conditioning image latent."})
|
||||||
|
}}
|
||||||
|
|
||||||
RETURN_TYPES = ("CONDITIONING", "CONDITIONING", "LATENT")
|
RETURN_TYPES = ("CONDITIONING", "CONDITIONING", "LATENT")
|
||||||
RETURN_NAMES = ("positive", "negative", "latent")
|
RETURN_NAMES = ("positive", "negative", "latent")
|
||||||
@ -40,12 +42,12 @@ class LTXVImgToVideo:
|
|||||||
CATEGORY = "conditioning/video_models"
|
CATEGORY = "conditioning/video_models"
|
||||||
FUNCTION = "generate"
|
FUNCTION = "generate"
|
||||||
|
|
||||||
def generate(self, positive, negative, image, vae, width, height, length, batch_size):
|
def generate(self, positive, negative, image, vae, width, height, length, batch_size, image_noise_scale):
|
||||||
pixels = comfy.utils.common_upscale(image.movedim(-1, 1), width, height, "bilinear", "center").movedim(1, -1)
|
pixels = comfy.utils.common_upscale(image.movedim(-1, 1), width, height, "bilinear", "center").movedim(1, -1)
|
||||||
encode_pixels = pixels[:, :, :, :3]
|
encode_pixels = pixels[:, :, :, :3]
|
||||||
t = vae.encode(encode_pixels)
|
t = vae.encode(encode_pixels)
|
||||||
positive = node_helpers.conditioning_set_values(positive, {"guiding_latent": t})
|
positive = node_helpers.conditioning_set_values(positive, {"guiding_latent": t, "guiding_latent_noise_scale": image_noise_scale})
|
||||||
negative = node_helpers.conditioning_set_values(negative, {"guiding_latent": t})
|
negative = node_helpers.conditioning_set_values(negative, {"guiding_latent": t, "guiding_latent_noise_scale": image_noise_scale})
|
||||||
|
|
||||||
latent = torch.zeros([batch_size, 128, ((length - 1) // 8) + 1, height // 32, width // 32], device=comfy.model_management.intermediate_device())
|
latent = torch.zeros([batch_size, 128, ((length - 1) // 8) + 1, height // 32, width // 32], device=comfy.model_management.intermediate_device())
|
||||||
latent[:, :, :t.shape[2]] = t
|
latent[:, :, :t.shape[2]] = t
|
||||||
@ -109,6 +111,7 @@ class ModelSamplingLTXV:
|
|||||||
model_sampling = ModelSamplingAdvanced(model.model.model_config)
|
model_sampling = ModelSamplingAdvanced(model.model.model_config)
|
||||||
model_sampling.set_parameters(shift=shift)
|
model_sampling.set_parameters(shift=shift)
|
||||||
m.add_object_patch("model_sampling", model_sampling)
|
m.add_object_patch("model_sampling", model_sampling)
|
||||||
|
|
||||||
return (m, )
|
return (m, )
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
41
comfy_extras/nodes_mahiro.py
Normal file
41
comfy_extras/nodes_mahiro.py
Normal file
@ -0,0 +1,41 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
class Mahiro:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {"model": ("MODEL",),
|
||||||
|
}}
|
||||||
|
RETURN_TYPES = ("MODEL",)
|
||||||
|
RETURN_NAMES = ("patched_model",)
|
||||||
|
FUNCTION = "patch"
|
||||||
|
CATEGORY = "_for_testing"
|
||||||
|
DESCRIPTION = "Modify the guidance to scale more on the 'direction' of the positive prompt rather than the difference between the negative prompt."
|
||||||
|
def patch(self, model):
|
||||||
|
m = model.clone()
|
||||||
|
def mahiro_normd(args):
|
||||||
|
scale: float = args['cond_scale']
|
||||||
|
cond_p: torch.Tensor = args['cond_denoised']
|
||||||
|
uncond_p: torch.Tensor = args['uncond_denoised']
|
||||||
|
#naive leap
|
||||||
|
leap = cond_p * scale
|
||||||
|
#sim with uncond leap
|
||||||
|
u_leap = uncond_p * scale
|
||||||
|
cfg = args["denoised"]
|
||||||
|
merge = (leap + cfg) / 2
|
||||||
|
normu = torch.sqrt(u_leap.abs()) * u_leap.sign()
|
||||||
|
normm = torch.sqrt(merge.abs()) * merge.sign()
|
||||||
|
sim = F.cosine_similarity(normu, normm).mean()
|
||||||
|
simsc = 2 * (sim+1)
|
||||||
|
wm = (simsc*cfg + (4-simsc)*leap) / 4
|
||||||
|
return wm
|
||||||
|
m.set_model_sampler_post_cfg_function(mahiro_normd)
|
||||||
|
return (m, )
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"Mahiro": Mahiro
|
||||||
|
}
|
||||||
|
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"Mahiro": "Mahiro is so cute that she deserves a better guidance function!! (。・ω・。)",
|
||||||
|
}
|
||||||
@ -1,4 +1,3 @@
|
|||||||
import folder_paths
|
|
||||||
import comfy.sd
|
import comfy.sd
|
||||||
import comfy.model_sampling
|
import comfy.model_sampling
|
||||||
import comfy.latent_formats
|
import comfy.latent_formats
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import torch
|
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
|
|
||||||
class PatchModelAddDownscale:
|
class PatchModelAddDownscale:
|
||||||
|
|||||||
@ -174,6 +174,28 @@ class ModelMergeMochiPreview(comfy_extras.nodes_model_merging.ModelMergeBlocks):
|
|||||||
|
|
||||||
return {"required": arg_dict}
|
return {"required": arg_dict}
|
||||||
|
|
||||||
|
class ModelMergeLTXV(comfy_extras.nodes_model_merging.ModelMergeBlocks):
|
||||||
|
CATEGORY = "advanced/model_merging/model_specific"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
arg_dict = { "model1": ("MODEL",),
|
||||||
|
"model2": ("MODEL",)}
|
||||||
|
|
||||||
|
argument = ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01})
|
||||||
|
|
||||||
|
arg_dict["patchify_proj."] = argument
|
||||||
|
arg_dict["adaln_single."] = argument
|
||||||
|
arg_dict["caption_projection."] = argument
|
||||||
|
|
||||||
|
for i in range(28):
|
||||||
|
arg_dict["transformer_blocks.{}.".format(i)] = argument
|
||||||
|
|
||||||
|
arg_dict["scale_shift_table"] = argument
|
||||||
|
arg_dict["proj_out."] = argument
|
||||||
|
|
||||||
|
return {"required": arg_dict}
|
||||||
|
|
||||||
NODE_CLASS_MAPPINGS = {
|
NODE_CLASS_MAPPINGS = {
|
||||||
"ModelMergeSD1": ModelMergeSD1,
|
"ModelMergeSD1": ModelMergeSD1,
|
||||||
"ModelMergeSD2": ModelMergeSD1, #SD1 and SD2 have the same blocks
|
"ModelMergeSD2": ModelMergeSD1, #SD1 and SD2 have the same blocks
|
||||||
@ -183,4 +205,5 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"ModelMergeFlux1": ModelMergeFlux1,
|
"ModelMergeFlux1": ModelMergeFlux1,
|
||||||
"ModelMergeSD35_Large": ModelMergeSD35_Large,
|
"ModelMergeSD35_Large": ModelMergeSD35_Large,
|
||||||
"ModelMergeMochiPreview": ModelMergeMochiPreview,
|
"ModelMergeMochiPreview": ModelMergeMochiPreview,
|
||||||
|
"ModelMergeLTXV": ModelMergeLTXV,
|
||||||
}
|
}
|
||||||
|
|||||||
@ -82,8 +82,7 @@ class CLIPTextEncodeSD3:
|
|||||||
tokens["l"] += empty["l"]
|
tokens["l"] += empty["l"]
|
||||||
while len(tokens["l"]) > len(tokens["g"]):
|
while len(tokens["l"]) > len(tokens["g"]):
|
||||||
tokens["g"] += empty["g"]
|
tokens["g"] += empty["g"]
|
||||||
cond, pooled = clip.encode_from_tokens(tokens, return_pooled=True)
|
return (clip.encode_from_tokens_scheduled(tokens), )
|
||||||
return ([[cond, {"pooled_output": pooled}]], )
|
|
||||||
|
|
||||||
|
|
||||||
class ControlNetApplySD3(nodes.ControlNetApplyAdvanced):
|
class ControlNetApplySD3(nodes.ControlNetApplyAdvanced):
|
||||||
|
|||||||
@ -16,7 +16,8 @@ class SkipLayerGuidanceDiT:
|
|||||||
"single_layers": ("STRING", {"default": "7, 8, 9", "multiline": False}),
|
"single_layers": ("STRING", {"default": "7, 8, 9", "multiline": False}),
|
||||||
"scale": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.1}),
|
"scale": ("FLOAT", {"default": 3.0, "min": 0.0, "max": 10.0, "step": 0.1}),
|
||||||
"start_percent": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 1.0, "step": 0.001}),
|
"start_percent": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
"end_percent": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001})
|
"end_percent": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}),
|
||||||
|
"rescaling_scale": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||||
}}
|
}}
|
||||||
RETURN_TYPES = ("MODEL",)
|
RETURN_TYPES = ("MODEL",)
|
||||||
FUNCTION = "skip_guidance"
|
FUNCTION = "skip_guidance"
|
||||||
@ -26,7 +27,7 @@ class SkipLayerGuidanceDiT:
|
|||||||
|
|
||||||
CATEGORY = "advanced/guidance"
|
CATEGORY = "advanced/guidance"
|
||||||
|
|
||||||
def skip_guidance(self, model, scale, start_percent, end_percent, double_layers="", single_layers=""):
|
def skip_guidance(self, model, scale, start_percent, end_percent, double_layers="", single_layers="", rescaling_scale=0):
|
||||||
# check if layer is comma separated integers
|
# check if layer is comma separated integers
|
||||||
def skip(args, extra_args):
|
def skip(args, extra_args):
|
||||||
return args
|
return args
|
||||||
@ -65,6 +66,11 @@ class SkipLayerGuidanceDiT:
|
|||||||
if scale > 0 and sigma_ >= sigma_end and sigma_ <= sigma_start:
|
if scale > 0 and sigma_ >= sigma_end and sigma_ <= sigma_start:
|
||||||
(slg,) = comfy.samplers.calc_cond_batch(model, [cond], x, sigma, model_options)
|
(slg,) = comfy.samplers.calc_cond_batch(model, [cond], x, sigma, model_options)
|
||||||
cfg_result = cfg_result + (cond_pred - slg) * scale
|
cfg_result = cfg_result + (cond_pred - slg) * scale
|
||||||
|
if rescaling_scale != 0:
|
||||||
|
factor = cond_pred.std() / cfg_result.std()
|
||||||
|
factor = rescaling_scale * factor + (1 - rescaling_scale)
|
||||||
|
cfg_result *= factor
|
||||||
|
|
||||||
return cfg_result
|
return cfg_result
|
||||||
|
|
||||||
m = model.clone()
|
m = model.clone()
|
||||||
|
|||||||
@ -1,4 +1,3 @@
|
|||||||
import os
|
|
||||||
import logging
|
import logging
|
||||||
from spandrel import ModelLoader, ImageModelDescriptor
|
from spandrel import ModelLoader, ImageModelDescriptor
|
||||||
from comfy import model_management
|
from comfy import model_management
|
||||||
|
|||||||
@ -1,7 +1,5 @@
|
|||||||
from PIL import Image, ImageOps
|
from PIL import Image
|
||||||
from io import BytesIO
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import struct
|
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import time
|
import time
|
||||||
|
|
||||||
|
|||||||
@ -16,7 +16,7 @@ import comfy.model_management
|
|||||||
from comfy_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker
|
from comfy_execution.graph import get_input_info, ExecutionList, DynamicPrompt, ExecutionBlocker
|
||||||
from comfy_execution.graph_utils import is_link, GraphBuilder
|
from comfy_execution.graph_utils import is_link, GraphBuilder
|
||||||
from comfy_execution.caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID
|
from comfy_execution.caching import HierarchicalCache, LRUCache, CacheKeySetInputSignature, CacheKeySetID
|
||||||
from comfy.cli_args import args
|
from comfy_execution.validation import validate_node_input
|
||||||
|
|
||||||
class ExecutionResult(Enum):
|
class ExecutionResult(Enum):
|
||||||
SUCCESS = 0
|
SUCCESS = 0
|
||||||
@ -480,7 +480,7 @@ class PromptExecutor:
|
|||||||
if self.caches.outputs.get(node_id) is not None:
|
if self.caches.outputs.get(node_id) is not None:
|
||||||
cached_nodes.append(node_id)
|
cached_nodes.append(node_id)
|
||||||
|
|
||||||
comfy.model_management.cleanup_models(keep_clone_weights_loaded=True)
|
comfy.model_management.cleanup_models_gc()
|
||||||
self.add_message("execution_cached",
|
self.add_message("execution_cached",
|
||||||
{ "nodes": cached_nodes, "prompt_id": prompt_id},
|
{ "nodes": cached_nodes, "prompt_id": prompt_id},
|
||||||
broadcast=False)
|
broadcast=False)
|
||||||
@ -527,7 +527,6 @@ class PromptExecutor:
|
|||||||
comfy.model_management.unload_all_models()
|
comfy.model_management.unload_all_models()
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
def validate_inputs(prompt, item, validated):
|
def validate_inputs(prompt, item, validated):
|
||||||
unique_id = item
|
unique_id = item
|
||||||
if unique_id in validated:
|
if unique_id in validated:
|
||||||
@ -589,8 +588,8 @@ def validate_inputs(prompt, item, validated):
|
|||||||
r = nodes.NODE_CLASS_MAPPINGS[o_class_type].RETURN_TYPES
|
r = nodes.NODE_CLASS_MAPPINGS[o_class_type].RETURN_TYPES
|
||||||
received_type = r[val[1]]
|
received_type = r[val[1]]
|
||||||
received_types[x] = received_type
|
received_types[x] = received_type
|
||||||
if 'input_types' not in validate_function_inputs and received_type != type_input:
|
if 'input_types' not in validate_function_inputs and not validate_node_input(received_type, type_input):
|
||||||
details = f"{x}, {received_type} != {type_input}"
|
details = f"{x}, received_type({received_type}) mismatch input_type({type_input})"
|
||||||
error = {
|
error = {
|
||||||
"type": "return_type_mismatch",
|
"type": "return_type_mismatch",
|
||||||
"message": "Return type mismatch between linked nodes",
|
"message": "Return type mismatch between linked nodes",
|
||||||
|
|||||||
36
fix_torch.py
36
fix_torch.py
@ -5,20 +5,24 @@ import ctypes
|
|||||||
import logging
|
import logging
|
||||||
|
|
||||||
|
|
||||||
torch_spec = importlib.util.find_spec("torch")
|
def fix_pytorch_libomp():
|
||||||
for folder in torch_spec.submodule_search_locations:
|
"""
|
||||||
lib_folder = os.path.join(folder, "lib")
|
Fix PyTorch libomp DLL issue on Windows by copying the correct DLL file if needed.
|
||||||
test_file = os.path.join(lib_folder, "fbgemm.dll")
|
"""
|
||||||
dest = os.path.join(lib_folder, "libomp140.x86_64.dll")
|
torch_spec = importlib.util.find_spec("torch")
|
||||||
if os.path.exists(dest):
|
for folder in torch_spec.submodule_search_locations:
|
||||||
break
|
lib_folder = os.path.join(folder, "lib")
|
||||||
|
test_file = os.path.join(lib_folder, "fbgemm.dll")
|
||||||
with open(test_file, 'rb') as f:
|
dest = os.path.join(lib_folder, "libomp140.x86_64.dll")
|
||||||
contents = f.read()
|
if os.path.exists(dest):
|
||||||
if b"libomp140.x86_64.dll" not in contents:
|
|
||||||
break
|
break
|
||||||
try:
|
|
||||||
mydll = ctypes.cdll.LoadLibrary(test_file)
|
with open(test_file, "rb") as f:
|
||||||
except FileNotFoundError as e:
|
contents = f.read()
|
||||||
logging.warning("Detected pytorch version with libomp issue, patching.")
|
if b"libomp140.x86_64.dll" not in contents:
|
||||||
shutil.copyfile(os.path.join(lib_folder, "libiomp5md.dll"), dest)
|
break
|
||||||
|
try:
|
||||||
|
mydll = ctypes.cdll.LoadLibrary(test_file)
|
||||||
|
except FileNotFoundError as e:
|
||||||
|
logging.warning("Detected pytorch version with libomp issue, patching.")
|
||||||
|
shutil.copyfile(os.path.join(lib_folder, "libiomp5md.dll"), dest)
|
||||||
|
|||||||
@ -4,7 +4,7 @@ import os
|
|||||||
import time
|
import time
|
||||||
import mimetypes
|
import mimetypes
|
||||||
import logging
|
import logging
|
||||||
from typing import Set, List, Dict, Tuple, Literal
|
from typing import Literal
|
||||||
from collections.abc import Collection
|
from collections.abc import Collection
|
||||||
|
|
||||||
supported_pt_extensions: set[str] = {'.ckpt', '.pt', '.bin', '.pth', '.safetensors', '.pkl', '.sft'}
|
supported_pt_extensions: set[str] = {'.ckpt', '.pt', '.bin', '.pth', '.safetensors', '.pkl', '.sft'}
|
||||||
@ -133,7 +133,7 @@ def get_directory_by_type(type_name: str) -> str | None:
|
|||||||
return get_input_directory()
|
return get_input_directory()
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def filter_files_content_types(files: List[str], content_types: Literal["image", "video", "audio"]) -> List[str]:
|
def filter_files_content_types(files: list[str], content_types: Literal["image", "video", "audio"]) -> list[str]:
|
||||||
"""
|
"""
|
||||||
Example:
|
Example:
|
||||||
files = os.listdir(folder_paths.get_input_directory())
|
files = os.listdir(folder_paths.get_input_directory())
|
||||||
|
|||||||
@ -1,7 +1,5 @@
|
|||||||
import torch
|
import torch
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
import struct
|
|
||||||
import numpy as np
|
|
||||||
from comfy.cli_args import args, LatentPreviewMethod
|
from comfy.cli_args import args, LatentPreviewMethod
|
||||||
from comfy.taesd.taesd import TAESD
|
from comfy.taesd.taesd import TAESD
|
||||||
import comfy.model_management
|
import comfy.model_management
|
||||||
|
|||||||
9
main.py
9
main.py
@ -8,6 +8,11 @@ import time
|
|||||||
from comfy.cli_args import args
|
from comfy.cli_args import args
|
||||||
from app.logger import setup_logger
|
from app.logger import setup_logger
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
#NOTE: These do not do anything on core ComfyUI which should already have no communication with the internet, they are for custom nodes.
|
||||||
|
os.environ['HF_HUB_DISABLE_TELEMETRY'] = '1'
|
||||||
|
os.environ['DO_NOT_TRACK'] = '1'
|
||||||
|
|
||||||
|
|
||||||
setup_logger(log_level=args.verbose)
|
setup_logger(log_level=args.verbose)
|
||||||
|
|
||||||
@ -82,7 +87,8 @@ if __name__ == "__main__":
|
|||||||
|
|
||||||
if args.windows_standalone_build:
|
if args.windows_standalone_build:
|
||||||
try:
|
try:
|
||||||
import fix_torch
|
from fix_torch import fix_pytorch_libomp
|
||||||
|
fix_pytorch_libomp()
|
||||||
except:
|
except:
|
||||||
pass
|
pass
|
||||||
|
|
||||||
@ -154,7 +160,6 @@ def prompt_worker(q, server):
|
|||||||
if need_gc:
|
if need_gc:
|
||||||
current_time = time.perf_counter()
|
current_time = time.perf_counter()
|
||||||
if (current_time - last_gc_collect) > gc_collect_interval:
|
if (current_time - last_gc_collect) > gc_collect_interval:
|
||||||
comfy.model_management.cleanup_models()
|
|
||||||
gc.collect()
|
gc.collect()
|
||||||
comfy.model_management.soft_empty_cache()
|
comfy.model_management.soft_empty_cache()
|
||||||
last_gc_collect = current_time
|
last_gc_collect = current_time
|
||||||
|
|||||||
@ -1,2 +0,0 @@
|
|||||||
# model_manager/__init__.py
|
|
||||||
from .download_models import download_model, DownloadModelStatus, DownloadStatusType, create_model_path, check_file_exists, track_download_progress, validate_filename
|
|
||||||
@ -1,234 +0,0 @@
|
|||||||
#NOTE: This was an experiment and WILL BE REMOVED
|
|
||||||
from __future__ import annotations
|
|
||||||
import aiohttp
|
|
||||||
import os
|
|
||||||
import traceback
|
|
||||||
import logging
|
|
||||||
from folder_paths import folder_names_and_paths, get_folder_paths
|
|
||||||
import re
|
|
||||||
from typing import Callable, Any, Optional, Awaitable, Dict
|
|
||||||
from enum import Enum
|
|
||||||
import time
|
|
||||||
from dataclasses import dataclass
|
|
||||||
|
|
||||||
|
|
||||||
class DownloadStatusType(Enum):
|
|
||||||
PENDING = "pending"
|
|
||||||
IN_PROGRESS = "in_progress"
|
|
||||||
COMPLETED = "completed"
|
|
||||||
ERROR = "error"
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class DownloadModelStatus():
|
|
||||||
status: str
|
|
||||||
progress_percentage: float
|
|
||||||
message: str
|
|
||||||
already_existed: bool = False
|
|
||||||
|
|
||||||
def __init__(self, status: DownloadStatusType, progress_percentage: float, message: str, already_existed: bool):
|
|
||||||
self.status = status.value # Store the string value of the Enum
|
|
||||||
self.progress_percentage = progress_percentage
|
|
||||||
self.message = message
|
|
||||||
self.already_existed = already_existed
|
|
||||||
|
|
||||||
def to_dict(self) -> Dict[str, Any]:
|
|
||||||
return {
|
|
||||||
"status": self.status,
|
|
||||||
"progress_percentage": self.progress_percentage,
|
|
||||||
"message": self.message,
|
|
||||||
"already_existed": self.already_existed
|
|
||||||
}
|
|
||||||
|
|
||||||
|
|
||||||
async def download_model(model_download_request: Callable[[str], Awaitable[aiohttp.ClientResponse]],
|
|
||||||
model_name: str,
|
|
||||||
model_url: str,
|
|
||||||
model_directory: str,
|
|
||||||
folder_path: str,
|
|
||||||
progress_callback: Callable[[str, DownloadModelStatus], Awaitable[Any]],
|
|
||||||
progress_interval: float = 1.0) -> DownloadModelStatus:
|
|
||||||
"""
|
|
||||||
Download a model file from a given URL into the models directory.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
model_download_request (Callable[[str], Awaitable[aiohttp.ClientResponse]]):
|
|
||||||
A function that makes an HTTP request. This makes it easier to mock in unit tests.
|
|
||||||
model_name (str):
|
|
||||||
The name of the model file to be downloaded. This will be the filename on disk.
|
|
||||||
model_url (str):
|
|
||||||
The URL from which to download the model.
|
|
||||||
model_directory (str):
|
|
||||||
The subdirectory within the main models directory where the model
|
|
||||||
should be saved (e.g., 'checkpoints', 'loras', etc.).
|
|
||||||
progress_callback (Callable[[str, DownloadModelStatus], Awaitable[Any]]):
|
|
||||||
An asynchronous function to call with progress updates.
|
|
||||||
folder_path (str);
|
|
||||||
Path to which model folder should be used as the root.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
DownloadModelStatus: The result of the download operation.
|
|
||||||
"""
|
|
||||||
if not validate_filename(model_name):
|
|
||||||
return DownloadModelStatus(
|
|
||||||
DownloadStatusType.ERROR,
|
|
||||||
0,
|
|
||||||
"Invalid model name",
|
|
||||||
False
|
|
||||||
)
|
|
||||||
|
|
||||||
if not model_directory in folder_names_and_paths:
|
|
||||||
return DownloadModelStatus(
|
|
||||||
DownloadStatusType.ERROR,
|
|
||||||
0,
|
|
||||||
"Invalid or unrecognized model directory. model_directory must be a known model type (eg 'checkpoints'). If you are seeing this error for a custom model type, ensure the relevant custom nodes are installed and working.",
|
|
||||||
False
|
|
||||||
)
|
|
||||||
|
|
||||||
if not folder_path in get_folder_paths(model_directory):
|
|
||||||
return DownloadModelStatus(
|
|
||||||
DownloadStatusType.ERROR,
|
|
||||||
0,
|
|
||||||
f"Invalid folder path '{folder_path}', does not match the list of known directories ({get_folder_paths(model_directory)}). If you're seeing this in the downloader UI, you may need to refresh the page.",
|
|
||||||
False
|
|
||||||
)
|
|
||||||
|
|
||||||
file_path = create_model_path(model_name, folder_path)
|
|
||||||
existing_file = await check_file_exists(file_path, model_name, progress_callback)
|
|
||||||
if existing_file:
|
|
||||||
return existing_file
|
|
||||||
|
|
||||||
try:
|
|
||||||
logging.info(f"Downloading {model_name} from {model_url}")
|
|
||||||
status = DownloadModelStatus(DownloadStatusType.PENDING, 0, f"Starting download of {model_name}", False)
|
|
||||||
await progress_callback(model_name, status)
|
|
||||||
|
|
||||||
response = await model_download_request(model_url)
|
|
||||||
if response.status != 200:
|
|
||||||
error_message = f"Failed to download {model_name}. Status code: {response.status}"
|
|
||||||
logging.error(error_message)
|
|
||||||
status = DownloadModelStatus(DownloadStatusType.ERROR, 0, error_message, False)
|
|
||||||
await progress_callback(model_name, status)
|
|
||||||
return DownloadModelStatus(DownloadStatusType.ERROR, 0, error_message, False)
|
|
||||||
|
|
||||||
return await track_download_progress(response, file_path, model_name, progress_callback, progress_interval)
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logging.error(f"Error in downloading model: {e}")
|
|
||||||
return await handle_download_error(e, model_name, progress_callback)
|
|
||||||
|
|
||||||
|
|
||||||
def create_model_path(model_name: str, folder_path: str) -> tuple[str, str]:
|
|
||||||
os.makedirs(folder_path, exist_ok=True)
|
|
||||||
file_path = os.path.join(folder_path, model_name)
|
|
||||||
|
|
||||||
# Ensure the resulting path is still within the base directory
|
|
||||||
abs_file_path = os.path.abspath(file_path)
|
|
||||||
abs_base_dir = os.path.abspath(folder_path)
|
|
||||||
if os.path.commonprefix([abs_file_path, abs_base_dir]) != abs_base_dir:
|
|
||||||
raise Exception(f"Invalid model directory: {folder_path}/{model_name}")
|
|
||||||
|
|
||||||
return file_path
|
|
||||||
|
|
||||||
|
|
||||||
async def check_file_exists(file_path: str,
|
|
||||||
model_name: str,
|
|
||||||
progress_callback: Callable[[str, DownloadModelStatus], Awaitable[Any]]
|
|
||||||
) -> Optional[DownloadModelStatus]:
|
|
||||||
if os.path.exists(file_path):
|
|
||||||
status = DownloadModelStatus(DownloadStatusType.COMPLETED, 100, f"{model_name} already exists", True)
|
|
||||||
await progress_callback(model_name, status)
|
|
||||||
return status
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
async def track_download_progress(response: aiohttp.ClientResponse,
|
|
||||||
file_path: str,
|
|
||||||
model_name: str,
|
|
||||||
progress_callback: Callable[[str, DownloadModelStatus], Awaitable[Any]],
|
|
||||||
interval: float = 1.0) -> DownloadModelStatus:
|
|
||||||
try:
|
|
||||||
total_size = int(response.headers.get('Content-Length', 0))
|
|
||||||
downloaded = 0
|
|
||||||
last_update_time = time.time()
|
|
||||||
|
|
||||||
async def update_progress():
|
|
||||||
nonlocal last_update_time
|
|
||||||
progress = (downloaded / total_size) * 100 if total_size > 0 else 0
|
|
||||||
status = DownloadModelStatus(DownloadStatusType.IN_PROGRESS, progress, f"Downloading {model_name}", False)
|
|
||||||
await progress_callback(model_name, status)
|
|
||||||
last_update_time = time.time()
|
|
||||||
|
|
||||||
temp_file_path = file_path + '.tmp'
|
|
||||||
with open(temp_file_path, 'wb') as f:
|
|
||||||
chunk_iterator = response.content.iter_chunked(8192)
|
|
||||||
while True:
|
|
||||||
try:
|
|
||||||
chunk = await chunk_iterator.__anext__()
|
|
||||||
except StopAsyncIteration:
|
|
||||||
break
|
|
||||||
f.write(chunk)
|
|
||||||
downloaded += len(chunk)
|
|
||||||
|
|
||||||
if time.time() - last_update_time >= interval:
|
|
||||||
await update_progress()
|
|
||||||
|
|
||||||
os.rename(temp_file_path, file_path)
|
|
||||||
|
|
||||||
await update_progress()
|
|
||||||
|
|
||||||
logging.info(f"Successfully downloaded {model_name}. Total downloaded: {downloaded}")
|
|
||||||
status = DownloadModelStatus(DownloadStatusType.COMPLETED, 100, f"Successfully downloaded {model_name}", False)
|
|
||||||
await progress_callback(model_name, status)
|
|
||||||
|
|
||||||
return status
|
|
||||||
except Exception as e:
|
|
||||||
logging.error(f"Error in track_download_progress: {e}")
|
|
||||||
logging.error(traceback.format_exc())
|
|
||||||
return await handle_download_error(e, model_name, progress_callback)
|
|
||||||
|
|
||||||
|
|
||||||
async def handle_download_error(e: Exception,
|
|
||||||
model_name: str,
|
|
||||||
progress_callback: Callable[[str, DownloadModelStatus], Any]
|
|
||||||
) -> DownloadModelStatus:
|
|
||||||
error_message = f"Error downloading {model_name}: {str(e)}"
|
|
||||||
status = DownloadModelStatus(DownloadStatusType.ERROR, 0, error_message, False)
|
|
||||||
await progress_callback(model_name, status)
|
|
||||||
return status
|
|
||||||
|
|
||||||
|
|
||||||
def validate_filename(filename: str)-> bool:
|
|
||||||
"""
|
|
||||||
Validate a filename to ensure it's safe and doesn't contain any path traversal attempts.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
filename (str): The filename to validate
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
bool: True if the filename is valid, False otherwise
|
|
||||||
"""
|
|
||||||
if not filename.lower().endswith(('.sft', '.safetensors')):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check if the filename is empty, None, or just whitespace
|
|
||||||
if not filename or not filename.strip():
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check for any directory traversal attempts or invalid characters
|
|
||||||
if any(char in filename for char in ['..', '/', '\\', '\n', '\r', '\t', '\0']):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Check if the filename starts with a dot (hidden file)
|
|
||||||
if filename.startswith('.'):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Use a whitelist of allowed characters
|
|
||||||
if not re.match(r'^[a-zA-Z0-9_\-. ]+$', filename):
|
|
||||||
return False
|
|
||||||
|
|
||||||
# Ensure the filename isn't too long
|
|
||||||
if len(filename) > 255:
|
|
||||||
return False
|
|
||||||
|
|
||||||
return True
|
|
||||||
42
nodes.py
42
nodes.py
@ -1,3 +1,4 @@
|
|||||||
|
from __future__ import annotations
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
import os
|
import os
|
||||||
@ -10,7 +11,7 @@ import time
|
|||||||
import random
|
import random
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
from PIL import Image, ImageOps, ImageSequence, ImageFile
|
from PIL import Image, ImageOps, ImageSequence
|
||||||
from PIL.PngImagePlugin import PngInfo
|
from PIL.PngImagePlugin import PngInfo
|
||||||
|
|
||||||
import numpy as np
|
import numpy as np
|
||||||
@ -24,6 +25,7 @@ import comfy.sample
|
|||||||
import comfy.sd
|
import comfy.sd
|
||||||
import comfy.utils
|
import comfy.utils
|
||||||
import comfy.controlnet
|
import comfy.controlnet
|
||||||
|
from comfy.comfy_types import IO, ComfyNodeABC, InputTypeDict
|
||||||
|
|
||||||
import comfy.clip_vision
|
import comfy.clip_vision
|
||||||
|
|
||||||
@ -44,16 +46,16 @@ def interrupt_processing(value=True):
|
|||||||
|
|
||||||
MAX_RESOLUTION=16384
|
MAX_RESOLUTION=16384
|
||||||
|
|
||||||
class CLIPTextEncode:
|
class CLIPTextEncode(ComfyNodeABC):
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s) -> InputTypeDict:
|
||||||
return {
|
return {
|
||||||
"required": {
|
"required": {
|
||||||
"text": ("STRING", {"multiline": True, "dynamicPrompts": True, "tooltip": "The text to be encoded."}),
|
"text": (IO.STRING, {"multiline": True, "dynamicPrompts": True, "tooltip": "The text to be encoded."}),
|
||||||
"clip": ("CLIP", {"tooltip": "The CLIP model used for encoding the text."})
|
"clip": (IO.CLIP, {"tooltip": "The CLIP model used for encoding the text."})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
RETURN_TYPES = ("CONDITIONING",)
|
RETURN_TYPES = (IO.CONDITIONING,)
|
||||||
OUTPUT_TOOLTIPS = ("A conditioning containing the embedded text used to guide the diffusion model.",)
|
OUTPUT_TOOLTIPS = ("A conditioning containing the embedded text used to guide the diffusion model.",)
|
||||||
FUNCTION = "encode"
|
FUNCTION = "encode"
|
||||||
|
|
||||||
@ -62,9 +64,8 @@ class CLIPTextEncode:
|
|||||||
|
|
||||||
def encode(self, clip, text):
|
def encode(self, clip, text):
|
||||||
tokens = clip.tokenize(text)
|
tokens = clip.tokenize(text)
|
||||||
output = clip.encode_from_tokens(tokens, return_pooled=True, return_dict=True)
|
return (clip.encode_from_tokens_scheduled(tokens), )
|
||||||
cond = output.pop("cond")
|
|
||||||
return ([[cond, output]], )
|
|
||||||
|
|
||||||
class ConditioningCombine:
|
class ConditioningCombine:
|
||||||
@classmethod
|
@classmethod
|
||||||
@ -392,7 +393,7 @@ class InpaintModelConditioning:
|
|||||||
|
|
||||||
CATEGORY = "conditioning/inpaint"
|
CATEGORY = "conditioning/inpaint"
|
||||||
|
|
||||||
def encode(self, positive, negative, pixels, vae, mask, noise_mask):
|
def encode(self, positive, negative, pixels, vae, mask, noise_mask=True):
|
||||||
x = (pixels.shape[1] // 8) * 8
|
x = (pixels.shape[1] // 8) * 8
|
||||||
y = (pixels.shape[2] // 8) * 8
|
y = (pixels.shape[2] // 8) * 8
|
||||||
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(pixels.shape[1], pixels.shape[2]), mode="bilinear")
|
mask = torch.nn.functional.interpolate(mask.reshape((-1, 1, mask.shape[-2], mask.shape[-1])), size=(pixels.shape[1], pixels.shape[2]), mode="bilinear")
|
||||||
@ -643,9 +644,7 @@ class LoraLoader:
|
|||||||
if self.loaded_lora[0] == lora_path:
|
if self.loaded_lora[0] == lora_path:
|
||||||
lora = self.loaded_lora[1]
|
lora = self.loaded_lora[1]
|
||||||
else:
|
else:
|
||||||
temp = self.loaded_lora
|
|
||||||
self.loaded_lora = None
|
self.loaded_lora = None
|
||||||
del temp
|
|
||||||
|
|
||||||
if lora is None:
|
if lora is None:
|
||||||
lora = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
lora = comfy.utils.load_torch_file(lora_path, safe_load=True)
|
||||||
@ -971,15 +970,19 @@ class CLIPVisionEncode:
|
|||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
return {"required": { "clip_vision": ("CLIP_VISION",),
|
return {"required": { "clip_vision": ("CLIP_VISION",),
|
||||||
"image": ("IMAGE",)
|
"image": ("IMAGE",),
|
||||||
|
"crop": (["center", "none"],)
|
||||||
}}
|
}}
|
||||||
RETURN_TYPES = ("CLIP_VISION_OUTPUT",)
|
RETURN_TYPES = ("CLIP_VISION_OUTPUT",)
|
||||||
FUNCTION = "encode"
|
FUNCTION = "encode"
|
||||||
|
|
||||||
CATEGORY = "conditioning"
|
CATEGORY = "conditioning"
|
||||||
|
|
||||||
def encode(self, clip_vision, image):
|
def encode(self, clip_vision, image, crop):
|
||||||
output = clip_vision.encode_image(image)
|
crop_image = True
|
||||||
|
if crop != "center":
|
||||||
|
crop_image = False
|
||||||
|
output = clip_vision.encode_image(image, crop=crop_image)
|
||||||
return (output,)
|
return (output,)
|
||||||
|
|
||||||
class StyleModelLoader:
|
class StyleModelLoader:
|
||||||
@ -1004,14 +1007,19 @@ class StyleModelApply:
|
|||||||
return {"required": {"conditioning": ("CONDITIONING", ),
|
return {"required": {"conditioning": ("CONDITIONING", ),
|
||||||
"style_model": ("STYLE_MODEL", ),
|
"style_model": ("STYLE_MODEL", ),
|
||||||
"clip_vision_output": ("CLIP_VISION_OUTPUT", ),
|
"clip_vision_output": ("CLIP_VISION_OUTPUT", ),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
|
||||||
|
"strength_type": (["multiply"], ),
|
||||||
}}
|
}}
|
||||||
RETURN_TYPES = ("CONDITIONING",)
|
RETURN_TYPES = ("CONDITIONING",)
|
||||||
FUNCTION = "apply_stylemodel"
|
FUNCTION = "apply_stylemodel"
|
||||||
|
|
||||||
CATEGORY = "conditioning/style_model"
|
CATEGORY = "conditioning/style_model"
|
||||||
|
|
||||||
def apply_stylemodel(self, clip_vision_output, style_model, conditioning):
|
def apply_stylemodel(self, clip_vision_output, style_model, conditioning, strength, strength_type):
|
||||||
cond = style_model.get_cond(clip_vision_output).flatten(start_dim=0, end_dim=1).unsqueeze(dim=0)
|
cond = style_model.get_cond(clip_vision_output).flatten(start_dim=0, end_dim=1).unsqueeze(dim=0)
|
||||||
|
if strength_type == "multiply":
|
||||||
|
cond *= strength
|
||||||
|
|
||||||
c = []
|
c = []
|
||||||
for t in conditioning:
|
for t in conditioning:
|
||||||
n = [torch.cat((t[0], cond), dim=1), t[1].copy()]
|
n = [torch.cat((t[0], cond), dim=1), t[1].copy()]
|
||||||
@ -2139,7 +2147,9 @@ def init_builtin_extra_nodes():
|
|||||||
"nodes_torch_compile.py",
|
"nodes_torch_compile.py",
|
||||||
"nodes_mochi.py",
|
"nodes_mochi.py",
|
||||||
"nodes_slg.py",
|
"nodes_slg.py",
|
||||||
|
"nodes_mahiro.py",
|
||||||
"nodes_lt.py",
|
"nodes_lt.py",
|
||||||
|
"nodes_hooks.py",
|
||||||
]
|
]
|
||||||
|
|
||||||
import_failed = []
|
import_failed = []
|
||||||
|
|||||||
@ -1,329 +1,328 @@
|
|||||||
{
|
{
|
||||||
"cells": [
|
"cells": [
|
||||||
{
|
{
|
||||||
"cell_type": "markdown",
|
"cell_type": "markdown",
|
||||||
"metadata": {
|
"metadata": {
|
||||||
"id": "aaaaaaaaaa"
|
"id": "aaaaaaaaaa"
|
||||||
},
|
},
|
||||||
"source": [
|
"source": [
|
||||||
"Git clone the repo and install the requirements. (ignore the pip errors about protobuf)"
|
"Git clone the repo and install the requirements. (ignore the pip errors about protobuf)"
|
||||||
]
|
]
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {
|
|
||||||
"id": "bbbbbbbbbb"
|
|
||||||
},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"#@title Environment Setup\n",
|
|
||||||
"\n",
|
|
||||||
"from pathlib import Path\n",
|
|
||||||
"\n",
|
|
||||||
"OPTIONS = {}\n",
|
|
||||||
"\n",
|
|
||||||
"USE_GOOGLE_DRIVE = False #@param {type:\"boolean\"}\n",
|
|
||||||
"UPDATE_COMFY_UI = True #@param {type:\"boolean\"}\n",
|
|
||||||
"WORKSPACE = 'ComfyUI'\n",
|
|
||||||
"OPTIONS['USE_GOOGLE_DRIVE'] = USE_GOOGLE_DRIVE\n",
|
|
||||||
"OPTIONS['UPDATE_COMFY_UI'] = UPDATE_COMFY_UI\n",
|
|
||||||
"\n",
|
|
||||||
"if OPTIONS['USE_GOOGLE_DRIVE']:\n",
|
|
||||||
" !echo \"Mounting Google Drive...\"\n",
|
|
||||||
" %cd /\n",
|
|
||||||
" \n",
|
|
||||||
" from google.colab import drive\n",
|
|
||||||
" drive.mount('/content/drive')\n",
|
|
||||||
"\n",
|
|
||||||
" WORKSPACE = \"/content/drive/MyDrive/ComfyUI\"\n",
|
|
||||||
" %cd /content/drive/MyDrive\n",
|
|
||||||
"\n",
|
|
||||||
"![ ! -d $WORKSPACE ] && echo -= Initial setup ComfyUI =- && git clone https://github.com/comfyanonymous/ComfyUI\n",
|
|
||||||
"%cd $WORKSPACE\n",
|
|
||||||
"\n",
|
|
||||||
"if OPTIONS['UPDATE_COMFY_UI']:\n",
|
|
||||||
" !echo -= Updating ComfyUI =-\n",
|
|
||||||
" !git pull\n",
|
|
||||||
"\n",
|
|
||||||
"!echo -= Install dependencies =-\n",
|
|
||||||
"!pip install xformers!=0.0.18 -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cu121 --extra-index-url https://download.pytorch.org/whl/cu118 --extra-index-url https://download.pytorch.org/whl/cu117"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "markdown",
|
|
||||||
"metadata": {
|
|
||||||
"id": "cccccccccc"
|
|
||||||
},
|
|
||||||
"source": [
|
|
||||||
"Download some models/checkpoints/vae or custom comfyui nodes (uncomment the commands for the ones you want)"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {
|
|
||||||
"id": "dddddddddd"
|
|
||||||
},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"# Checkpoints\n",
|
|
||||||
"\n",
|
|
||||||
"### SDXL\n",
|
|
||||||
"### I recommend these workflow examples: https://comfyanonymous.github.io/ComfyUI_examples/sdxl/\n",
|
|
||||||
"\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_base_1.0.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-xl-refiner-1.0/resolve/main/sd_xl_refiner_1.0.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"\n",
|
|
||||||
"# SDXL ReVision\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/clip_vision_g/resolve/main/clip_vision_g.safetensors -P ./models/clip_vision/\n",
|
|
||||||
"\n",
|
|
||||||
"# SD1.5\n",
|
|
||||||
"!wget -c https://huggingface.co/Comfy-Org/stable-diffusion-v1-5-archive/resolve/main/v1-5-pruned-emaonly-fp16.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"\n",
|
|
||||||
"# SD2\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-2-1-base/resolve/main/v2-1_512-ema-pruned.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-2-1/resolve/main/v2-1_768-ema-pruned.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"\n",
|
|
||||||
"# Some SD1.5 anime style\n",
|
|
||||||
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/Models/AbyssOrangeMix2/AbyssOrangeMix2_hard.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/Models/AbyssOrangeMix3/AOM3A1_orangemixs.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/Models/AbyssOrangeMix3/AOM3A3_orangemixs.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"#!wget -c https://huggingface.co/Linaqruf/anything-v3.0/resolve/main/anything-v3-fp16-pruned.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"\n",
|
|
||||||
"# Waifu Diffusion 1.5 (anime style SD2.x 768-v)\n",
|
|
||||||
"#!wget -c https://huggingface.co/waifu-diffusion/wd-1-5-beta3/resolve/main/wd-illusion-fp16.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# unCLIP models\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/illuminatiDiffusionV1_v11_unCLIP/resolve/main/illuminatiDiffusionV1_v11-unclip-h-fp16.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/wd-1.5-beta2_unCLIP/resolve/main/wd-1-5-beta2-aesthetic-unclip-h-fp16.safetensors -P ./models/checkpoints/\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# VAE\n",
|
|
||||||
"!wget -c https://huggingface.co/stabilityai/sd-vae-ft-mse-original/resolve/main/vae-ft-mse-840000-ema-pruned.safetensors -P ./models/vae/\n",
|
|
||||||
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/VAEs/orangemix.vae.pt -P ./models/vae/\n",
|
|
||||||
"#!wget -c https://huggingface.co/hakurei/waifu-diffusion-v1-4/resolve/main/vae/kl-f8-anime2.ckpt -P ./models/vae/\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# Loras\n",
|
|
||||||
"#!wget -c https://civitai.com/api/download/models/10350 -O ./models/loras/theovercomer8sContrastFix_sd21768.safetensors #theovercomer8sContrastFix SD2.x 768-v\n",
|
|
||||||
"#!wget -c https://civitai.com/api/download/models/10638 -O ./models/loras/theovercomer8sContrastFix_sd15.safetensors #theovercomer8sContrastFix SD1.x\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_offset_example-lora_1.0.safetensors -P ./models/loras/ #SDXL offset noise lora\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# T2I-Adapter\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_depth_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_seg_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_sketch_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_keypose_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_openpose_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_color_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_canny_sd14v1.pth -P ./models/controlnet/\n",
|
|
||||||
"\n",
|
|
||||||
"# T2I Styles Model\n",
|
|
||||||
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_style_sd14v1.pth -P ./models/style_models/\n",
|
|
||||||
"\n",
|
|
||||||
"# CLIPVision model (needed for styles model)\n",
|
|
||||||
"#!wget -c https://huggingface.co/openai/clip-vit-large-patch14/resolve/main/pytorch_model.bin -O ./models/clip_vision/clip_vit14.bin\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# ControlNet\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11e_sd15_ip2p_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11e_sd15_shuffle_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_canny_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11f1p_sd15_depth_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_inpaint_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_lineart_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_mlsd_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_normalbae_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_openpose_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_scribble_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_seg_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_softedge_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15s2_lineart_anime_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11u_sd15_tile_fp16.safetensors -P ./models/controlnet/\n",
|
|
||||||
"\n",
|
|
||||||
"# ControlNet SDXL\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-canny-rank256.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-depth-rank256.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-recolor-rank256.safetensors -P ./models/controlnet/\n",
|
|
||||||
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-sketch-rank256.safetensors -P ./models/controlnet/\n",
|
|
||||||
"\n",
|
|
||||||
"# Controlnet Preprocessor nodes by Fannovel16\n",
|
|
||||||
"#!cd custom_nodes && git clone https://github.com/Fannovel16/comfy_controlnet_preprocessors; cd comfy_controlnet_preprocessors && python install.py\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# GLIGEN\n",
|
|
||||||
"#!wget -c https://huggingface.co/comfyanonymous/GLIGEN_pruned_safetensors/resolve/main/gligen_sd14_textbox_pruned_fp16.safetensors -P ./models/gligen/\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"# ESRGAN upscale model\n",
|
|
||||||
"#!wget -c https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth -P ./models/upscale_models/\n",
|
|
||||||
"#!wget -c https://huggingface.co/sberbank-ai/Real-ESRGAN/resolve/main/RealESRGAN_x2.pth -P ./models/upscale_models/\n",
|
|
||||||
"#!wget -c https://huggingface.co/sberbank-ai/Real-ESRGAN/resolve/main/RealESRGAN_x4.pth -P ./models/upscale_models/\n",
|
|
||||||
"\n",
|
|
||||||
"\n"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "markdown",
|
|
||||||
"metadata": {
|
|
||||||
"id": "kkkkkkkkkkkkkkk"
|
|
||||||
},
|
|
||||||
"source": [
|
|
||||||
"### Run ComfyUI with cloudflared (Recommended Way)\n",
|
|
||||||
"\n",
|
|
||||||
"\n"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {
|
|
||||||
"id": "jjjjjjjjjjjjjj"
|
|
||||||
},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"!wget https://github.com/cloudflare/cloudflared/releases/latest/download/cloudflared-linux-amd64.deb\n",
|
|
||||||
"!dpkg -i cloudflared-linux-amd64.deb\n",
|
|
||||||
"\n",
|
|
||||||
"import subprocess\n",
|
|
||||||
"import threading\n",
|
|
||||||
"import time\n",
|
|
||||||
"import socket\n",
|
|
||||||
"import urllib.request\n",
|
|
||||||
"\n",
|
|
||||||
"def iframe_thread(port):\n",
|
|
||||||
" while True:\n",
|
|
||||||
" time.sleep(0.5)\n",
|
|
||||||
" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n",
|
|
||||||
" result = sock.connect_ex(('127.0.0.1', port))\n",
|
|
||||||
" if result == 0:\n",
|
|
||||||
" break\n",
|
|
||||||
" sock.close()\n",
|
|
||||||
" print(\"\\nComfyUI finished loading, trying to launch cloudflared (if it gets stuck here cloudflared is having issues)\\n\")\n",
|
|
||||||
"\n",
|
|
||||||
" p = subprocess.Popen([\"cloudflared\", \"tunnel\", \"--url\", \"http://127.0.0.1:{}\".format(port)], stdout=subprocess.PIPE, stderr=subprocess.PIPE)\n",
|
|
||||||
" for line in p.stderr:\n",
|
|
||||||
" l = line.decode()\n",
|
|
||||||
" if \"trycloudflare.com \" in l:\n",
|
|
||||||
" print(\"This is the URL to access ComfyUI:\", l[l.find(\"http\"):], end='')\n",
|
|
||||||
" #print(l, end='')\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"threading.Thread(target=iframe_thread, daemon=True, args=(8188,)).start()\n",
|
|
||||||
"\n",
|
|
||||||
"!python main.py --dont-print-server"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "markdown",
|
|
||||||
"metadata": {
|
|
||||||
"id": "kkkkkkkkkkkkkk"
|
|
||||||
},
|
|
||||||
"source": [
|
|
||||||
"### Run ComfyUI with localtunnel\n",
|
|
||||||
"\n",
|
|
||||||
"\n"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {
|
|
||||||
"id": "jjjjjjjjjjjjj"
|
|
||||||
},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"!npm install -g localtunnel\n",
|
|
||||||
"\n",
|
|
||||||
"import subprocess\n",
|
|
||||||
"import threading\n",
|
|
||||||
"import time\n",
|
|
||||||
"import socket\n",
|
|
||||||
"import urllib.request\n",
|
|
||||||
"\n",
|
|
||||||
"def iframe_thread(port):\n",
|
|
||||||
" while True:\n",
|
|
||||||
" time.sleep(0.5)\n",
|
|
||||||
" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n",
|
|
||||||
" result = sock.connect_ex(('127.0.0.1', port))\n",
|
|
||||||
" if result == 0:\n",
|
|
||||||
" break\n",
|
|
||||||
" sock.close()\n",
|
|
||||||
" print(\"\\nComfyUI finished loading, trying to launch localtunnel (if it gets stuck here localtunnel is having issues)\\n\")\n",
|
|
||||||
"\n",
|
|
||||||
" print(\"The password/enpoint ip for localtunnel is:\", urllib.request.urlopen('https://ipv4.icanhazip.com').read().decode('utf8').strip(\"\\n\"))\n",
|
|
||||||
" p = subprocess.Popen([\"lt\", \"--port\", \"{}\".format(port)], stdout=subprocess.PIPE)\n",
|
|
||||||
" for line in p.stdout:\n",
|
|
||||||
" print(line.decode(), end='')\n",
|
|
||||||
"\n",
|
|
||||||
"\n",
|
|
||||||
"threading.Thread(target=iframe_thread, daemon=True, args=(8188,)).start()\n",
|
|
||||||
"\n",
|
|
||||||
"!python main.py --dont-print-server"
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "markdown",
|
|
||||||
"metadata": {
|
|
||||||
"id": "gggggggggg"
|
|
||||||
},
|
|
||||||
"source": [
|
|
||||||
"### Run ComfyUI with colab iframe (use only in case the previous way with localtunnel doesn't work)\n",
|
|
||||||
"\n",
|
|
||||||
"You should see the ui appear in an iframe. If you get a 403 error, it's your firefox settings or an extension that's messing things up.\n",
|
|
||||||
"\n",
|
|
||||||
"If you want to open it in another window use the link.\n",
|
|
||||||
"\n",
|
|
||||||
"Note that some UI features like live image previews won't work because the colab iframe blocks websockets."
|
|
||||||
]
|
|
||||||
},
|
|
||||||
{
|
|
||||||
"cell_type": "code",
|
|
||||||
"execution_count": null,
|
|
||||||
"metadata": {
|
|
||||||
"id": "hhhhhhhhhh"
|
|
||||||
},
|
|
||||||
"outputs": [],
|
|
||||||
"source": [
|
|
||||||
"import threading\n",
|
|
||||||
"import time\n",
|
|
||||||
"import socket\n",
|
|
||||||
"def iframe_thread(port):\n",
|
|
||||||
" while True:\n",
|
|
||||||
" time.sleep(0.5)\n",
|
|
||||||
" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n",
|
|
||||||
" result = sock.connect_ex(('127.0.0.1', port))\n",
|
|
||||||
" if result == 0:\n",
|
|
||||||
" break\n",
|
|
||||||
" sock.close()\n",
|
|
||||||
" from google.colab import output\n",
|
|
||||||
" output.serve_kernel_port_as_iframe(port, height=1024)\n",
|
|
||||||
" print(\"to open it in a window you can open this link here:\")\n",
|
|
||||||
" output.serve_kernel_port_as_window(port)\n",
|
|
||||||
"\n",
|
|
||||||
"threading.Thread(target=iframe_thread, daemon=True, args=(8188,)).start()\n",
|
|
||||||
"\n",
|
|
||||||
"!python main.py --dont-print-server"
|
|
||||||
]
|
|
||||||
}
|
|
||||||
],
|
|
||||||
"metadata": {
|
|
||||||
"accelerator": "GPU",
|
|
||||||
"colab": {
|
|
||||||
"provenance": []
|
|
||||||
},
|
|
||||||
"gpuClass": "standard",
|
|
||||||
"kernelspec": {
|
|
||||||
"display_name": "Python 3",
|
|
||||||
"name": "python3"
|
|
||||||
},
|
|
||||||
"language_info": {
|
|
||||||
"name": "python"
|
|
||||||
}
|
|
||||||
},
|
},
|
||||||
"nbformat": 4,
|
{
|
||||||
"nbformat_minor": 0
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "bbbbbbbbbb"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"#@title Environment Setup\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"OPTIONS = {}\n",
|
||||||
|
"\n",
|
||||||
|
"USE_GOOGLE_DRIVE = False #@param {type:\"boolean\"}\n",
|
||||||
|
"UPDATE_COMFY_UI = True #@param {type:\"boolean\"}\n",
|
||||||
|
"WORKSPACE = 'ComfyUI'\n",
|
||||||
|
"OPTIONS['USE_GOOGLE_DRIVE'] = USE_GOOGLE_DRIVE\n",
|
||||||
|
"OPTIONS['UPDATE_COMFY_UI'] = UPDATE_COMFY_UI\n",
|
||||||
|
"\n",
|
||||||
|
"if OPTIONS['USE_GOOGLE_DRIVE']:\n",
|
||||||
|
" !echo \"Mounting Google Drive...\"\n",
|
||||||
|
" %cd /\n",
|
||||||
|
" \n",
|
||||||
|
" from google.colab import drive\n",
|
||||||
|
" drive.mount('/content/drive')\n",
|
||||||
|
"\n",
|
||||||
|
" WORKSPACE = \"/content/drive/MyDrive/ComfyUI\"\n",
|
||||||
|
" %cd /content/drive/MyDrive\n",
|
||||||
|
"\n",
|
||||||
|
"![ ! -d $WORKSPACE ] && echo -= Initial setup ComfyUI =- && git clone https://github.com/comfyanonymous/ComfyUI\n",
|
||||||
|
"%cd $WORKSPACE\n",
|
||||||
|
"\n",
|
||||||
|
"if OPTIONS['UPDATE_COMFY_UI']:\n",
|
||||||
|
" !echo -= Updating ComfyUI =-\n",
|
||||||
|
" !git pull\n",
|
||||||
|
"\n",
|
||||||
|
"!echo -= Install dependencies =-\n",
|
||||||
|
"!pip install xformers!=0.0.18 -r requirements.txt --extra-index-url https://download.pytorch.org/whl/cu121 --extra-index-url https://download.pytorch.org/whl/cu118 --extra-index-url https://download.pytorch.org/whl/cu117"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"metadata": {
|
||||||
|
"id": "cccccccccc"
|
||||||
|
},
|
||||||
|
"source": [
|
||||||
|
"Download some models/checkpoints/vae or custom comfyui nodes (uncomment the commands for the ones you want)"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "dddddddddd"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"# Checkpoints\n",
|
||||||
|
"\n",
|
||||||
|
"### SDXL\n",
|
||||||
|
"### I recommend these workflow examples: https://comfyanonymous.github.io/ComfyUI_examples/sdxl/\n",
|
||||||
|
"\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_base_1.0.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-xl-refiner-1.0/resolve/main/sd_xl_refiner_1.0.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"\n",
|
||||||
|
"# SDXL ReVision\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/clip_vision_g/resolve/main/clip_vision_g.safetensors -P ./models/clip_vision/\n",
|
||||||
|
"\n",
|
||||||
|
"# SD1.5\n",
|
||||||
|
"!wget -c https://huggingface.co/Comfy-Org/stable-diffusion-v1-5-archive/resolve/main/v1-5-pruned-emaonly-fp16.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"\n",
|
||||||
|
"# SD2\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-2-1-base/resolve/main/v2-1_512-ema-pruned.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-2-1/resolve/main/v2-1_768-ema-pruned.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"\n",
|
||||||
|
"# Some SD1.5 anime style\n",
|
||||||
|
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/Models/AbyssOrangeMix2/AbyssOrangeMix2_hard.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/Models/AbyssOrangeMix3/AOM3A1_orangemixs.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/Models/AbyssOrangeMix3/AOM3A3_orangemixs.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"#!wget -c https://huggingface.co/Linaqruf/anything-v3.0/resolve/main/anything-v3-fp16-pruned.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"\n",
|
||||||
|
"# Waifu Diffusion 1.5 (anime style SD2.x 768-v)\n",
|
||||||
|
"#!wget -c https://huggingface.co/waifu-diffusion/wd-1-5-beta3/resolve/main/wd-illusion-fp16.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# unCLIP models\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/illuminatiDiffusionV1_v11_unCLIP/resolve/main/illuminatiDiffusionV1_v11-unclip-h-fp16.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/wd-1.5-beta2_unCLIP/resolve/main/wd-1-5-beta2-aesthetic-unclip-h-fp16.safetensors -P ./models/checkpoints/\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# VAE\n",
|
||||||
|
"!wget -c https://huggingface.co/stabilityai/sd-vae-ft-mse-original/resolve/main/vae-ft-mse-840000-ema-pruned.safetensors -P ./models/vae/\n",
|
||||||
|
"#!wget -c https://huggingface.co/WarriorMama777/OrangeMixs/resolve/main/VAEs/orangemix.vae.pt -P ./models/vae/\n",
|
||||||
|
"#!wget -c https://huggingface.co/hakurei/waifu-diffusion-v1-4/resolve/main/vae/kl-f8-anime2.ckpt -P ./models/vae/\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# Loras\n",
|
||||||
|
"#!wget -c https://civitai.com/api/download/models/10350 -O ./models/loras/theovercomer8sContrastFix_sd21768.safetensors #theovercomer8sContrastFix SD2.x 768-v\n",
|
||||||
|
"#!wget -c https://civitai.com/api/download/models/10638 -O ./models/loras/theovercomer8sContrastFix_sd15.safetensors #theovercomer8sContrastFix SD1.x\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/stable-diffusion-xl-base-1.0/resolve/main/sd_xl_offset_example-lora_1.0.safetensors -P ./models/loras/ #SDXL offset noise lora\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# T2I-Adapter\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_depth_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_seg_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_sketch_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_keypose_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_openpose_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_color_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_canny_sd14v1.pth -P ./models/controlnet/\n",
|
||||||
|
"\n",
|
||||||
|
"# T2I Styles Model\n",
|
||||||
|
"#!wget -c https://huggingface.co/TencentARC/T2I-Adapter/resolve/main/models/t2iadapter_style_sd14v1.pth -P ./models/style_models/\n",
|
||||||
|
"\n",
|
||||||
|
"# CLIPVision model (needed for styles model)\n",
|
||||||
|
"#!wget -c https://huggingface.co/openai/clip-vit-large-patch14/resolve/main/pytorch_model.bin -O ./models/clip_vision/clip_vit14.bin\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# ControlNet\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11e_sd15_ip2p_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11e_sd15_shuffle_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_canny_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11f1p_sd15_depth_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_inpaint_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_lineart_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_mlsd_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_normalbae_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_openpose_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_scribble_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_seg_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15_softedge_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11p_sd15s2_lineart_anime_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors/resolve/main/control_v11u_sd15_tile_fp16.safetensors -P ./models/controlnet/\n",
|
||||||
|
"\n",
|
||||||
|
"# ControlNet SDXL\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-canny-rank256.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-depth-rank256.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-recolor-rank256.safetensors -P ./models/controlnet/\n",
|
||||||
|
"#!wget -c https://huggingface.co/stabilityai/control-lora/resolve/main/control-LoRAs-rank256/control-lora-sketch-rank256.safetensors -P ./models/controlnet/\n",
|
||||||
|
"\n",
|
||||||
|
"# Controlnet Preprocessor nodes by Fannovel16\n",
|
||||||
|
"#!cd custom_nodes && git clone https://github.com/Fannovel16/comfy_controlnet_preprocessors; cd comfy_controlnet_preprocessors && python install.py\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# GLIGEN\n",
|
||||||
|
"#!wget -c https://huggingface.co/comfyanonymous/GLIGEN_pruned_safetensors/resolve/main/gligen_sd14_textbox_pruned_fp16.safetensors -P ./models/gligen/\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"# ESRGAN upscale model\n",
|
||||||
|
"#!wget -c https://github.com/xinntao/Real-ESRGAN/releases/download/v0.1.0/RealESRGAN_x4plus.pth -P ./models/upscale_models/\n",
|
||||||
|
"#!wget -c https://huggingface.co/sberbank-ai/Real-ESRGAN/resolve/main/RealESRGAN_x2.pth -P ./models/upscale_models/\n",
|
||||||
|
"#!wget -c https://huggingface.co/sberbank-ai/Real-ESRGAN/resolve/main/RealESRGAN_x4.pth -P ./models/upscale_models/\n",
|
||||||
|
"\n",
|
||||||
|
"\n"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"metadata": {
|
||||||
|
"id": "kkkkkkkkkkkkkkk"
|
||||||
|
},
|
||||||
|
"source": [
|
||||||
|
"### Run ComfyUI with cloudflared (Recommended Way)\n",
|
||||||
|
"\n",
|
||||||
|
"\n"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "jjjjjjjjjjjjjj"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"!wget https://github.com/cloudflare/cloudflared/releases/latest/download/cloudflared-linux-amd64.deb\n",
|
||||||
|
"!dpkg -i cloudflared-linux-amd64.deb\n",
|
||||||
|
"\n",
|
||||||
|
"import subprocess\n",
|
||||||
|
"import threading\n",
|
||||||
|
"import time\n",
|
||||||
|
"import socket\n",
|
||||||
|
"import urllib.request\n",
|
||||||
|
"\n",
|
||||||
|
"def iframe_thread(port):\n",
|
||||||
|
" while True:\n",
|
||||||
|
" time.sleep(0.5)\n",
|
||||||
|
" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n",
|
||||||
|
" result = sock.connect_ex(('127.0.0.1', port))\n",
|
||||||
|
" if result == 0:\n",
|
||||||
|
" break\n",
|
||||||
|
" sock.close()\n",
|
||||||
|
" print(\"\\nComfyUI finished loading, trying to launch cloudflared (if it gets stuck here cloudflared is having issues)\\n\")\n",
|
||||||
|
"\n",
|
||||||
|
" p = subprocess.Popen([\"cloudflared\", \"tunnel\", \"--url\", \"http://127.0.0.1:{}\".format(port)], stdout=subprocess.PIPE, stderr=subprocess.PIPE)\n",
|
||||||
|
" for line in p.stderr:\n",
|
||||||
|
" l = line.decode()\n",
|
||||||
|
" if \"trycloudflare.com \" in l:\n",
|
||||||
|
" print(\"This is the URL to access ComfyUI:\", l[l.find(\"http\"):], end='')\n",
|
||||||
|
" #print(l, end='')\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"threading.Thread(target=iframe_thread, daemon=True, args=(8188,)).start()\n",
|
||||||
|
"\n",
|
||||||
|
"!python main.py --dont-print-server"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"metadata": {
|
||||||
|
"id": "kkkkkkkkkkkkkk"
|
||||||
|
},
|
||||||
|
"source": [
|
||||||
|
"### Run ComfyUI with localtunnel\n",
|
||||||
|
"\n",
|
||||||
|
"\n"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "jjjjjjjjjjjjj"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"!npm install -g localtunnel\n",
|
||||||
|
"\n",
|
||||||
|
"import subprocess\n",
|
||||||
|
"import threading\n",
|
||||||
|
"import time\n",
|
||||||
|
"import socket\n",
|
||||||
|
"import urllib.request\n",
|
||||||
|
"\n",
|
||||||
|
"def iframe_thread(port):\n",
|
||||||
|
" while True:\n",
|
||||||
|
" time.sleep(0.5)\n",
|
||||||
|
" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n",
|
||||||
|
" result = sock.connect_ex(('127.0.0.1', port))\n",
|
||||||
|
" if result == 0:\n",
|
||||||
|
" break\n",
|
||||||
|
" sock.close()\n",
|
||||||
|
" print(\"\\nComfyUI finished loading, trying to launch localtunnel (if it gets stuck here localtunnel is having issues)\\n\")\n",
|
||||||
|
"\n",
|
||||||
|
" print(\"The password/enpoint ip for localtunnel is:\", urllib.request.urlopen('https://ipv4.icanhazip.com').read().decode('utf8').strip(\"\\n\"))\n",
|
||||||
|
" p = subprocess.Popen([\"lt\", \"--port\", \"{}\".format(port)], stdout=subprocess.PIPE)\n",
|
||||||
|
" for line in p.stdout:\n",
|
||||||
|
" print(line.decode(), end='')\n",
|
||||||
|
"\n",
|
||||||
|
"\n",
|
||||||
|
"threading.Thread(target=iframe_thread, daemon=True, args=(8188,)).start()\n",
|
||||||
|
"\n",
|
||||||
|
"!python main.py --dont-print-server"
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "markdown",
|
||||||
|
"metadata": {
|
||||||
|
"id": "gggggggggg"
|
||||||
|
},
|
||||||
|
"source": [
|
||||||
|
"### Run ComfyUI with colab iframe (use only in case the previous way with localtunnel doesn't work)\n",
|
||||||
|
"\n",
|
||||||
|
"You should see the ui appear in an iframe. If you get a 403 error, it's your firefox settings or an extension that's messing things up.\n",
|
||||||
|
"\n",
|
||||||
|
"If you want to open it in another window use the link.\n",
|
||||||
|
"\n",
|
||||||
|
"Note that some UI features like live image previews won't work because the colab iframe blocks websockets."
|
||||||
|
]
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"cell_type": "code",
|
||||||
|
"execution_count": null,
|
||||||
|
"metadata": {
|
||||||
|
"id": "hhhhhhhhhh"
|
||||||
|
},
|
||||||
|
"outputs": [],
|
||||||
|
"source": [
|
||||||
|
"import threading\n",
|
||||||
|
"import time\n",
|
||||||
|
"import socket\n",
|
||||||
|
"def iframe_thread(port):\n",
|
||||||
|
" while True:\n",
|
||||||
|
" time.sleep(0.5)\n",
|
||||||
|
" sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)\n",
|
||||||
|
" result = sock.connect_ex(('127.0.0.1', port))\n",
|
||||||
|
" if result == 0:\n",
|
||||||
|
" break\n",
|
||||||
|
" sock.close()\n",
|
||||||
|
" from google.colab import output\n",
|
||||||
|
" output.serve_kernel_port_as_iframe(port, height=1024)\n",
|
||||||
|
" print(\"to open it in a window you can open this link here:\")\n",
|
||||||
|
" output.serve_kernel_port_as_window(port)\n",
|
||||||
|
"\n",
|
||||||
|
"threading.Thread(target=iframe_thread, daemon=True, args=(8188,)).start()\n",
|
||||||
|
"\n",
|
||||||
|
"!python main.py --dont-print-server"
|
||||||
|
]
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"metadata": {
|
||||||
|
"accelerator": "GPU",
|
||||||
|
"colab": {
|
||||||
|
"provenance": []
|
||||||
|
},
|
||||||
|
"gpuClass": "standard",
|
||||||
|
"kernelspec": {
|
||||||
|
"display_name": "Python 3",
|
||||||
|
"name": "python3"
|
||||||
|
},
|
||||||
|
"language_info": {
|
||||||
|
"name": "python"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"nbformat": 4,
|
||||||
|
"nbformat_minor": 0
|
||||||
}
|
}
|
||||||
|
|||||||
8
ruff.toml
Normal file
8
ruff.toml
Normal file
@ -0,0 +1,8 @@
|
|||||||
|
# Disable all rules by default
|
||||||
|
lint.ignore = ["ALL"]
|
||||||
|
|
||||||
|
# Enable specific rules
|
||||||
|
lint.select = [
|
||||||
|
"S307", # suspicious-eval-usage
|
||||||
|
"F401", # unused-import
|
||||||
|
]
|
||||||
@ -1,6 +1,5 @@
|
|||||||
import json
|
import json
|
||||||
from urllib import request, parse
|
from urllib import request
|
||||||
import random
|
|
||||||
|
|
||||||
#This is the ComfyUI api prompt format.
|
#This is the ComfyUI api prompt format.
|
||||||
|
|
||||||
|
|||||||
31
server.py
31
server.py
@ -30,7 +30,6 @@ import node_helpers
|
|||||||
from app.frontend_management import FrontendManager
|
from app.frontend_management import FrontendManager
|
||||||
from app.user_manager import UserManager
|
from app.user_manager import UserManager
|
||||||
from app.model_manager import ModelFileManager
|
from app.model_manager import ModelFileManager
|
||||||
from model_filemanager import download_model, DownloadModelStatus
|
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
from api_server.routes.internal.internal_routes import InternalRoutes
|
from api_server.routes.internal.internal_routes import InternalRoutes
|
||||||
|
|
||||||
@ -678,36 +677,6 @@ class PromptServer():
|
|||||||
self.prompt_queue.delete_history_item(id_to_delete)
|
self.prompt_queue.delete_history_item(id_to_delete)
|
||||||
|
|
||||||
return web.Response(status=200)
|
return web.Response(status=200)
|
||||||
|
|
||||||
# Internal route. Should not be depended upon and is subject to change at any time.
|
|
||||||
# TODO(robinhuang): Move to internal route table class once we refactor PromptServer to pass around Websocket.
|
|
||||||
# NOTE: This was an experiment and WILL BE REMOVED
|
|
||||||
@routes.post("/internal/models/download")
|
|
||||||
async def download_handler(request):
|
|
||||||
async def report_progress(filename: str, status: DownloadModelStatus):
|
|
||||||
payload = status.to_dict()
|
|
||||||
payload['download_path'] = filename
|
|
||||||
await self.send_json("download_progress", payload)
|
|
||||||
|
|
||||||
data = await request.json()
|
|
||||||
url = data.get('url')
|
|
||||||
model_directory = data.get('model_directory')
|
|
||||||
folder_path = data.get('folder_path')
|
|
||||||
model_filename = data.get('model_filename')
|
|
||||||
progress_interval = data.get('progress_interval', 1.0) # In seconds, how often to report download progress.
|
|
||||||
|
|
||||||
if not url or not model_directory or not model_filename or not folder_path:
|
|
||||||
return web.json_response({"status": "error", "message": "Missing URL or folder path or filename"}, status=400)
|
|
||||||
|
|
||||||
session = self.client_session
|
|
||||||
if session is None:
|
|
||||||
logging.error("Client session is not initialized")
|
|
||||||
return web.Response(status=500)
|
|
||||||
|
|
||||||
task = asyncio.create_task(download_model(lambda url: session.get(url), model_filename, url, model_directory, folder_path, report_progress, progress_interval))
|
|
||||||
await task
|
|
||||||
|
|
||||||
return web.json_response(task.result().to_dict())
|
|
||||||
|
|
||||||
async def setup(self):
|
async def setup(self):
|
||||||
timeout = aiohttp.ClientTimeout(total=None) # no timeout
|
timeout = aiohttp.ClientTimeout(total=None) # no timeout
|
||||||
|
|||||||
119
tests-unit/execution_test/validate_node_input_test.py
Normal file
119
tests-unit/execution_test/validate_node_input_test.py
Normal file
@ -0,0 +1,119 @@
|
|||||||
|
import pytest
|
||||||
|
from comfy_execution.validation import validate_node_input
|
||||||
|
|
||||||
|
|
||||||
|
def test_exact_match():
|
||||||
|
"""Test cases where types match exactly"""
|
||||||
|
assert validate_node_input("STRING", "STRING")
|
||||||
|
assert validate_node_input("STRING,INT", "STRING,INT")
|
||||||
|
assert validate_node_input("INT,STRING", "STRING,INT") # Order shouldn't matter
|
||||||
|
|
||||||
|
|
||||||
|
def test_strict_mode():
|
||||||
|
"""Test strict mode validation"""
|
||||||
|
# Should pass - received type is subset of input type
|
||||||
|
assert validate_node_input("STRING", "STRING,INT", strict=True)
|
||||||
|
assert validate_node_input("INT", "STRING,INT", strict=True)
|
||||||
|
assert validate_node_input("STRING,INT", "STRING,INT,BOOLEAN", strict=True)
|
||||||
|
|
||||||
|
# Should fail - received type is not subset of input type
|
||||||
|
assert not validate_node_input("STRING,INT", "STRING", strict=True)
|
||||||
|
assert not validate_node_input("STRING,BOOLEAN", "STRING", strict=True)
|
||||||
|
assert not validate_node_input("INT,BOOLEAN", "STRING,INT", strict=True)
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_strict_mode():
|
||||||
|
"""Test non-strict mode validation (default behavior)"""
|
||||||
|
# Should pass - types have overlap
|
||||||
|
assert validate_node_input("STRING,BOOLEAN", "STRING,INT")
|
||||||
|
assert validate_node_input("STRING,INT", "INT,BOOLEAN")
|
||||||
|
assert validate_node_input("STRING", "STRING,INT")
|
||||||
|
|
||||||
|
# Should fail - no overlap in types
|
||||||
|
assert not validate_node_input("BOOLEAN", "STRING,INT")
|
||||||
|
assert not validate_node_input("FLOAT", "STRING,INT")
|
||||||
|
assert not validate_node_input("FLOAT,BOOLEAN", "STRING,INT")
|
||||||
|
|
||||||
|
|
||||||
|
def test_whitespace_handling():
|
||||||
|
"""Test that whitespace is handled correctly"""
|
||||||
|
assert validate_node_input("STRING, INT", "STRING,INT")
|
||||||
|
assert validate_node_input("STRING,INT", "STRING, INT")
|
||||||
|
assert validate_node_input(" STRING , INT ", "STRING,INT")
|
||||||
|
assert validate_node_input("STRING,INT", " STRING , INT ")
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty_strings():
|
||||||
|
"""Test behavior with empty strings"""
|
||||||
|
assert validate_node_input("", "")
|
||||||
|
assert not validate_node_input("STRING", "")
|
||||||
|
assert not validate_node_input("", "STRING")
|
||||||
|
|
||||||
|
|
||||||
|
def test_single_vs_multiple():
|
||||||
|
"""Test single type against multiple types"""
|
||||||
|
assert validate_node_input("STRING", "STRING,INT,BOOLEAN")
|
||||||
|
assert validate_node_input("STRING,INT,BOOLEAN", "STRING", strict=False)
|
||||||
|
assert not validate_node_input("STRING,INT,BOOLEAN", "STRING", strict=True)
|
||||||
|
|
||||||
|
|
||||||
|
def test_non_string():
|
||||||
|
"""Test non-string types"""
|
||||||
|
obj1 = object()
|
||||||
|
obj2 = object()
|
||||||
|
assert validate_node_input(obj1, obj1)
|
||||||
|
assert not validate_node_input(obj1, obj2)
|
||||||
|
|
||||||
|
|
||||||
|
class NotEqualsOverrideTest(str):
|
||||||
|
"""Test class for ``__ne__`` override."""
|
||||||
|
|
||||||
|
def __ne__(self, value: object) -> bool:
|
||||||
|
if self == "*" or value == "*":
|
||||||
|
return False
|
||||||
|
if self == "LONGER_THAN_2":
|
||||||
|
return not len(value) > 2
|
||||||
|
raise TypeError("This is a class for unit tests only.")
|
||||||
|
|
||||||
|
|
||||||
|
def test_ne_override():
|
||||||
|
"""Test ``__ne__`` any override"""
|
||||||
|
any = NotEqualsOverrideTest("*")
|
||||||
|
invalid_type = "INVALID_TYPE"
|
||||||
|
obj = object()
|
||||||
|
assert validate_node_input(any, any)
|
||||||
|
assert validate_node_input(any, invalid_type)
|
||||||
|
assert validate_node_input(any, obj)
|
||||||
|
assert validate_node_input(any, {})
|
||||||
|
assert validate_node_input(any, [])
|
||||||
|
assert validate_node_input(any, [1, 2, 3])
|
||||||
|
|
||||||
|
|
||||||
|
def test_ne_custom_override():
|
||||||
|
"""Test ``__ne__`` custom override"""
|
||||||
|
special = NotEqualsOverrideTest("LONGER_THAN_2")
|
||||||
|
|
||||||
|
assert validate_node_input(special, special)
|
||||||
|
assert validate_node_input(special, "*")
|
||||||
|
assert validate_node_input(special, "INVALID_TYPE")
|
||||||
|
assert validate_node_input(special, [1, 2, 3])
|
||||||
|
|
||||||
|
# Should fail
|
||||||
|
assert not validate_node_input(special, [1, 2])
|
||||||
|
assert not validate_node_input(special, "TY")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"received,input_type,strict,expected",
|
||||||
|
[
|
||||||
|
("STRING", "STRING", False, True),
|
||||||
|
("STRING,INT", "STRING,INT", False, True),
|
||||||
|
("STRING", "STRING,INT", True, True),
|
||||||
|
("STRING,INT", "STRING", True, False),
|
||||||
|
("BOOLEAN", "STRING,INT", False, False),
|
||||||
|
("STRING,BOOLEAN", "STRING,INT", False, True),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
def test_parametrized_cases(received, input_type, strict, expected):
|
||||||
|
"""Parametrized test cases for various scenarios"""
|
||||||
|
assert validate_node_input(received, input_type, strict) == expected
|
||||||
@ -1,337 +0,0 @@
|
|||||||
import pytest
|
|
||||||
import tempfile
|
|
||||||
import aiohttp
|
|
||||||
from aiohttp import ClientResponse
|
|
||||||
import itertools
|
|
||||||
import os
|
|
||||||
from unittest.mock import AsyncMock, patch, MagicMock
|
|
||||||
from model_filemanager import download_model, track_download_progress, create_model_path, check_file_exists, DownloadStatusType, DownloadModelStatus, validate_filename
|
|
||||||
import folder_paths
|
|
||||||
|
|
||||||
@pytest.fixture
|
|
||||||
def temp_dir():
|
|
||||||
with tempfile.TemporaryDirectory() as tmpdirname:
|
|
||||||
yield tmpdirname
|
|
||||||
|
|
||||||
class AsyncIteratorMock:
|
|
||||||
"""
|
|
||||||
A mock class that simulates an asynchronous iterator.
|
|
||||||
This is used to mimic the behavior of aiohttp's content iterator.
|
|
||||||
"""
|
|
||||||
def __init__(self, seq):
|
|
||||||
# Convert the input sequence into an iterator
|
|
||||||
self.iter = iter(seq)
|
|
||||||
|
|
||||||
def __aiter__(self):
|
|
||||||
# This method is called when 'async for' is used
|
|
||||||
return self
|
|
||||||
|
|
||||||
async def __anext__(self):
|
|
||||||
# This method is called for each iteration in an 'async for' loop
|
|
||||||
try:
|
|
||||||
return next(self.iter)
|
|
||||||
except StopIteration:
|
|
||||||
# This is the asynchronous equivalent of StopIteration
|
|
||||||
raise StopAsyncIteration
|
|
||||||
|
|
||||||
class ContentMock:
|
|
||||||
"""
|
|
||||||
A mock class that simulates the content attribute of an aiohttp ClientResponse.
|
|
||||||
This class provides the iter_chunked method which returns an async iterator of chunks.
|
|
||||||
"""
|
|
||||||
def __init__(self, chunks):
|
|
||||||
# Store the chunks that will be returned by the iterator
|
|
||||||
self.chunks = chunks
|
|
||||||
|
|
||||||
def iter_chunked(self, chunk_size):
|
|
||||||
# This method mimics aiohttp's content.iter_chunked()
|
|
||||||
# For simplicity in testing, we ignore chunk_size and just return our predefined chunks
|
|
||||||
return AsyncIteratorMock(self.chunks)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_download_model_success(temp_dir):
|
|
||||||
mock_response = AsyncMock(spec=aiohttp.ClientResponse)
|
|
||||||
mock_response.status = 200
|
|
||||||
mock_response.headers = {'Content-Length': '1000'}
|
|
||||||
# Create a mock for content that returns an async iterator directly
|
|
||||||
chunks = [b'a' * 500, b'b' * 300, b'c' * 200]
|
|
||||||
mock_response.content = ContentMock(chunks)
|
|
||||||
|
|
||||||
mock_make_request = AsyncMock(return_value=mock_response)
|
|
||||||
mock_progress_callback = AsyncMock()
|
|
||||||
|
|
||||||
time_values = itertools.count(0, 0.1)
|
|
||||||
|
|
||||||
fake_paths = {'checkpoints': ([temp_dir], folder_paths.supported_pt_extensions)}
|
|
||||||
|
|
||||||
with patch('model_filemanager.create_model_path', return_value=('models/checkpoints/model.sft', 'model.sft')), \
|
|
||||||
patch('model_filemanager.check_file_exists', return_value=None), \
|
|
||||||
patch('folder_paths.folder_names_and_paths', fake_paths), \
|
|
||||||
patch('time.time', side_effect=time_values): # Simulate time passing
|
|
||||||
|
|
||||||
result = await download_model(
|
|
||||||
mock_make_request,
|
|
||||||
'model.sft',
|
|
||||||
'http://example.com/model.sft',
|
|
||||||
'checkpoints',
|
|
||||||
temp_dir,
|
|
||||||
mock_progress_callback
|
|
||||||
)
|
|
||||||
|
|
||||||
# Assert the result
|
|
||||||
assert isinstance(result, DownloadModelStatus)
|
|
||||||
assert result.message == 'Successfully downloaded model.sft'
|
|
||||||
assert result.status == 'completed'
|
|
||||||
assert result.already_existed is False
|
|
||||||
|
|
||||||
# Check progress callback calls
|
|
||||||
assert mock_progress_callback.call_count >= 3 # At least start, one progress update, and completion
|
|
||||||
|
|
||||||
# Check initial call
|
|
||||||
mock_progress_callback.assert_any_call(
|
|
||||||
'model.sft',
|
|
||||||
DownloadModelStatus(DownloadStatusType.PENDING, 0, "Starting download of model.sft", False)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Check final call
|
|
||||||
mock_progress_callback.assert_any_call(
|
|
||||||
'model.sft',
|
|
||||||
DownloadModelStatus(DownloadStatusType.COMPLETED, 100, "Successfully downloaded model.sft", False)
|
|
||||||
)
|
|
||||||
|
|
||||||
mock_file_path = os.path.join(temp_dir, 'model.sft')
|
|
||||||
assert os.path.exists(mock_file_path)
|
|
||||||
with open(mock_file_path, 'rb') as mock_file:
|
|
||||||
assert mock_file.read() == b''.join(chunks)
|
|
||||||
os.remove(mock_file_path)
|
|
||||||
|
|
||||||
# Verify request was made
|
|
||||||
mock_make_request.assert_called_once_with('http://example.com/model.sft')
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_download_model_url_request_failure(temp_dir):
|
|
||||||
# Mock dependencies
|
|
||||||
mock_response = AsyncMock(spec=ClientResponse)
|
|
||||||
mock_response.status = 404 # Simulate a "Not Found" error
|
|
||||||
mock_get = AsyncMock(return_value=mock_response)
|
|
||||||
mock_progress_callback = AsyncMock()
|
|
||||||
|
|
||||||
fake_paths = {'checkpoints': ([temp_dir], folder_paths.supported_pt_extensions)}
|
|
||||||
|
|
||||||
# Mock the create_model_path function
|
|
||||||
with patch('model_filemanager.create_model_path', return_value='/mock/path/model.safetensors'), \
|
|
||||||
patch('model_filemanager.check_file_exists', return_value=None), \
|
|
||||||
patch('folder_paths.folder_names_and_paths', fake_paths):
|
|
||||||
# Call the function
|
|
||||||
result = await download_model(
|
|
||||||
mock_get,
|
|
||||||
'model.safetensors',
|
|
||||||
'http://example.com/model.safetensors',
|
|
||||||
'checkpoints',
|
|
||||||
temp_dir,
|
|
||||||
mock_progress_callback
|
|
||||||
)
|
|
||||||
|
|
||||||
# Assert the expected behavior
|
|
||||||
assert isinstance(result, DownloadModelStatus)
|
|
||||||
assert result.status == 'error'
|
|
||||||
assert result.message == 'Failed to download model.safetensors. Status code: 404'
|
|
||||||
assert result.already_existed is False
|
|
||||||
|
|
||||||
# Check that progress_callback was called with the correct arguments
|
|
||||||
mock_progress_callback.assert_any_call(
|
|
||||||
'model.safetensors',
|
|
||||||
DownloadModelStatus(
|
|
||||||
status=DownloadStatusType.PENDING,
|
|
||||||
progress_percentage=0,
|
|
||||||
message='Starting download of model.safetensors',
|
|
||||||
already_existed=False
|
|
||||||
)
|
|
||||||
)
|
|
||||||
mock_progress_callback.assert_called_with(
|
|
||||||
'model.safetensors',
|
|
||||||
DownloadModelStatus(
|
|
||||||
status=DownloadStatusType.ERROR,
|
|
||||||
progress_percentage=0,
|
|
||||||
message='Failed to download model.safetensors. Status code: 404',
|
|
||||||
already_existed=False
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Verify that the get method was called with the correct URL
|
|
||||||
mock_get.assert_called_once_with('http://example.com/model.safetensors')
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_download_model_invalid_model_subdirectory():
|
|
||||||
mock_make_request = AsyncMock()
|
|
||||||
mock_progress_callback = AsyncMock()
|
|
||||||
|
|
||||||
result = await download_model(
|
|
||||||
mock_make_request,
|
|
||||||
'model.sft',
|
|
||||||
'http://example.com/model.sft',
|
|
||||||
'../bad_path',
|
|
||||||
'../bad_path',
|
|
||||||
mock_progress_callback
|
|
||||||
)
|
|
||||||
|
|
||||||
# Assert the result
|
|
||||||
assert isinstance(result, DownloadModelStatus)
|
|
||||||
assert result.message.startswith('Invalid or unrecognized model directory')
|
|
||||||
assert result.status == 'error'
|
|
||||||
assert result.already_existed is False
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_download_model_invalid_folder_path():
|
|
||||||
mock_make_request = AsyncMock()
|
|
||||||
mock_progress_callback = AsyncMock()
|
|
||||||
|
|
||||||
result = await download_model(
|
|
||||||
mock_make_request,
|
|
||||||
'model.sft',
|
|
||||||
'http://example.com/model.sft',
|
|
||||||
'checkpoints',
|
|
||||||
'invalid_path',
|
|
||||||
mock_progress_callback
|
|
||||||
)
|
|
||||||
|
|
||||||
# Assert the result
|
|
||||||
assert isinstance(result, DownloadModelStatus)
|
|
||||||
assert result.message.startswith("Invalid folder path")
|
|
||||||
assert result.status == 'error'
|
|
||||||
assert result.already_existed is False
|
|
||||||
|
|
||||||
def test_create_model_path(tmp_path, monkeypatch):
|
|
||||||
model_name = "model.safetensors"
|
|
||||||
folder_path = os.path.join(tmp_path, "mock_dir")
|
|
||||||
|
|
||||||
file_path = create_model_path(model_name, folder_path)
|
|
||||||
|
|
||||||
assert file_path == os.path.join(folder_path, "model.safetensors")
|
|
||||||
assert os.path.exists(os.path.dirname(file_path))
|
|
||||||
|
|
||||||
with pytest.raises(Exception, match="Invalid model directory"):
|
|
||||||
create_model_path("../path_traversal.safetensors", folder_path)
|
|
||||||
|
|
||||||
with pytest.raises(Exception, match="Invalid model directory"):
|
|
||||||
create_model_path("/etc/some_root_path", folder_path)
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_check_file_exists_when_file_exists(tmp_path):
|
|
||||||
file_path = tmp_path / "existing_model.sft"
|
|
||||||
file_path.touch() # Create an empty file
|
|
||||||
|
|
||||||
mock_callback = AsyncMock()
|
|
||||||
|
|
||||||
result = await check_file_exists(str(file_path), "existing_model.sft", mock_callback)
|
|
||||||
|
|
||||||
assert result is not None
|
|
||||||
assert result.status == "completed"
|
|
||||||
assert result.message == "existing_model.sft already exists"
|
|
||||||
assert result.already_existed is True
|
|
||||||
|
|
||||||
mock_callback.assert_called_once_with(
|
|
||||||
"existing_model.sft",
|
|
||||||
DownloadModelStatus(DownloadStatusType.COMPLETED, 100, "existing_model.sft already exists", already_existed=True)
|
|
||||||
)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_check_file_exists_when_file_does_not_exist(tmp_path):
|
|
||||||
file_path = tmp_path / "non_existing_model.sft"
|
|
||||||
|
|
||||||
mock_callback = AsyncMock()
|
|
||||||
|
|
||||||
result = await check_file_exists(str(file_path), "non_existing_model.sft", mock_callback)
|
|
||||||
|
|
||||||
assert result is None
|
|
||||||
mock_callback.assert_not_called()
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_track_download_progress_no_content_length(temp_dir):
|
|
||||||
mock_response = AsyncMock(spec=aiohttp.ClientResponse)
|
|
||||||
mock_response.headers = {} # No Content-Length header
|
|
||||||
chunks = [b'a' * 500, b'b' * 500]
|
|
||||||
mock_response.content.iter_chunked.return_value = AsyncIteratorMock(chunks)
|
|
||||||
|
|
||||||
mock_callback = AsyncMock()
|
|
||||||
|
|
||||||
full_path = os.path.join(temp_dir, 'model.sft')
|
|
||||||
|
|
||||||
result = await track_download_progress(
|
|
||||||
mock_response, full_path, 'model.sft',
|
|
||||||
mock_callback, interval=0.1
|
|
||||||
)
|
|
||||||
|
|
||||||
assert result.status == "completed"
|
|
||||||
|
|
||||||
assert os.path.exists(full_path)
|
|
||||||
with open(full_path, 'rb') as f:
|
|
||||||
assert f.read() == b''.join(chunks)
|
|
||||||
os.remove(full_path)
|
|
||||||
|
|
||||||
# Check that progress was reported even without knowing the total size
|
|
||||||
mock_callback.assert_any_call(
|
|
||||||
'model.sft',
|
|
||||||
DownloadModelStatus(DownloadStatusType.IN_PROGRESS, 0, "Downloading model.sft", already_existed=False)
|
|
||||||
)
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_track_download_progress_interval(temp_dir):
|
|
||||||
mock_response = AsyncMock(spec=aiohttp.ClientResponse)
|
|
||||||
mock_response.headers = {'Content-Length': '1000'}
|
|
||||||
chunks = [b'a' * 100] * 10
|
|
||||||
mock_response.content.iter_chunked.return_value = AsyncIteratorMock(chunks)
|
|
||||||
|
|
||||||
mock_callback = AsyncMock()
|
|
||||||
mock_open = MagicMock(return_value=MagicMock())
|
|
||||||
|
|
||||||
# Create a mock time function that returns incremental float values
|
|
||||||
mock_time = MagicMock()
|
|
||||||
mock_time.side_effect = [i * 0.5 for i in range(30)] # This should be enough for 10 chunks
|
|
||||||
|
|
||||||
full_path = os.path.join(temp_dir, 'model.sft')
|
|
||||||
|
|
||||||
with patch('time.time', mock_time):
|
|
||||||
await track_download_progress(
|
|
||||||
mock_response, full_path, 'model.sft',
|
|
||||||
mock_callback, interval=1.0
|
|
||||||
)
|
|
||||||
|
|
||||||
assert os.path.exists(full_path)
|
|
||||||
with open(full_path, 'rb') as f:
|
|
||||||
assert f.read() == b''.join(chunks)
|
|
||||||
os.remove(full_path)
|
|
||||||
|
|
||||||
# Assert that progress was updated at least 3 times (start, at least one interval, and end)
|
|
||||||
assert mock_callback.call_count >= 3, f"Expected at least 3 calls, but got {mock_callback.call_count}"
|
|
||||||
|
|
||||||
# Verify the first and last calls
|
|
||||||
first_call = mock_callback.call_args_list[0]
|
|
||||||
assert first_call[0][1].status == "in_progress"
|
|
||||||
# Allow for some initial progress, but it should be less than 50%
|
|
||||||
assert 0 <= first_call[0][1].progress_percentage < 50, f"First call progress was {first_call[0][1].progress_percentage}%"
|
|
||||||
|
|
||||||
last_call = mock_callback.call_args_list[-1]
|
|
||||||
assert last_call[0][1].status == "completed"
|
|
||||||
assert last_call[0][1].progress_percentage == 100
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("filename, expected", [
|
|
||||||
("valid_model.safetensors", True),
|
|
||||||
("valid_model.sft", True),
|
|
||||||
("valid model.safetensors", True), # Test with space
|
|
||||||
("UPPERCASE_MODEL.SAFETENSORS", True),
|
|
||||||
("model_with.multiple.dots.pt", False),
|
|
||||||
("", False), # Empty string
|
|
||||||
("../../../etc/passwd", False), # Path traversal attempt
|
|
||||||
("/etc/passwd", False), # Absolute path
|
|
||||||
("\\windows\\system32\\config\\sam", False), # Windows path
|
|
||||||
(".hidden_file.pt", False), # Hidden file
|
|
||||||
("invalid<char>.ckpt", False), # Invalid character
|
|
||||||
("invalid?.ckpt", False), # Another invalid character
|
|
||||||
("very" * 100 + ".safetensors", False), # Too long filename
|
|
||||||
("\nmodel_with_newline.pt", False), # Newline character
|
|
||||||
("model_with_emoji😊.pt", False), # Emoji in filename
|
|
||||||
])
|
|
||||||
def test_validate_filename(filename, expected):
|
|
||||||
assert validate_filename(filename) == expected
|
|
||||||
@ -1,6 +1,5 @@
|
|||||||
from copy import deepcopy
|
from copy import deepcopy
|
||||||
from io import BytesIO
|
from io import BytesIO
|
||||||
from urllib import request
|
|
||||||
import numpy
|
import numpy
|
||||||
import os
|
import os
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|||||||
@ -24,5 +24,8 @@ def load_extra_path_config(yaml_path):
|
|||||||
full_path = y
|
full_path = y
|
||||||
if base_path is not None:
|
if base_path is not None:
|
||||||
full_path = os.path.join(base_path, full_path)
|
full_path = os.path.join(base_path, full_path)
|
||||||
|
elif not os.path.isabs(full_path):
|
||||||
|
yaml_dir = os.path.dirname(os.path.abspath(yaml_path))
|
||||||
|
full_path = os.path.abspath(os.path.join(yaml_dir, y))
|
||||||
logging.info("Adding extra search path {} {}".format(x, full_path))
|
logging.info("Adding extra search path {} {}".format(x, full_path))
|
||||||
folder_paths.add_model_folder_path(x, full_path, is_default)
|
folder_paths.add_model_folder_path(x, full_path, is_default)
|
||||||
|
|||||||
103
web/assets/ExtensionPanel-CfMfcLgI.js
generated
vendored
103
web/assets/ExtensionPanel-CfMfcLgI.js
generated
vendored
@ -1,103 +0,0 @@
|
|||||||
var __defProp = Object.defineProperty;
|
|
||||||
var __name = (target, value) => __defProp(target, "name", { value, configurable: true });
|
|
||||||
import { d as defineComponent, c6 as useExtensionStore, u as useSettingStore, r as ref, o as onMounted, q as computed, g as openBlock, h as createElementBlock, i as createVNode, y as withCtx, z as unref, bT as script$1, A as createBaseVNode, x as createBlock, N as Fragment, O as renderList, a6 as toDisplayString, aw as createTextVNode, bR as script$3, j as createCommentVNode, D as script$4 } from "./index-B6dYHNhg.js";
|
|
||||||
import { s as script, a as script$2 } from "./index-CjwCGacA.js";
|
|
||||||
import "./index-MX9DEi8Q.js";
|
|
||||||
const _hoisted_1 = { class: "extension-panel" };
|
|
||||||
const _hoisted_2 = { class: "mt-4" };
|
|
||||||
const _sfc_main = /* @__PURE__ */ defineComponent({
|
|
||||||
__name: "ExtensionPanel",
|
|
||||||
setup(__props) {
|
|
||||||
const extensionStore = useExtensionStore();
|
|
||||||
const settingStore = useSettingStore();
|
|
||||||
const editingEnabledExtensions = ref({});
|
|
||||||
onMounted(() => {
|
|
||||||
extensionStore.extensions.forEach((ext) => {
|
|
||||||
editingEnabledExtensions.value[ext.name] = extensionStore.isExtensionEnabled(ext.name);
|
|
||||||
});
|
|
||||||
});
|
|
||||||
const changedExtensions = computed(() => {
|
|
||||||
return extensionStore.extensions.filter(
|
|
||||||
(ext) => editingEnabledExtensions.value[ext.name] !== extensionStore.isExtensionEnabled(ext.name)
|
|
||||||
);
|
|
||||||
});
|
|
||||||
const hasChanges = computed(() => {
|
|
||||||
return changedExtensions.value.length > 0;
|
|
||||||
});
|
|
||||||
const updateExtensionStatus = /* @__PURE__ */ __name(() => {
|
|
||||||
const editingDisabledExtensionNames = Object.entries(
|
|
||||||
editingEnabledExtensions.value
|
|
||||||
).filter(([_, enabled]) => !enabled).map(([name]) => name);
|
|
||||||
settingStore.set("Comfy.Extension.Disabled", [
|
|
||||||
...extensionStore.inactiveDisabledExtensionNames,
|
|
||||||
...editingDisabledExtensionNames
|
|
||||||
]);
|
|
||||||
}, "updateExtensionStatus");
|
|
||||||
const applyChanges = /* @__PURE__ */ __name(() => {
|
|
||||||
window.location.reload();
|
|
||||||
}, "applyChanges");
|
|
||||||
return (_ctx, _cache) => {
|
|
||||||
return openBlock(), createElementBlock("div", _hoisted_1, [
|
|
||||||
createVNode(unref(script$2), {
|
|
||||||
value: unref(extensionStore).extensions,
|
|
||||||
stripedRows: "",
|
|
||||||
size: "small"
|
|
||||||
}, {
|
|
||||||
default: withCtx(() => [
|
|
||||||
createVNode(unref(script), {
|
|
||||||
field: "name",
|
|
||||||
header: _ctx.$t("extensionName"),
|
|
||||||
sortable: ""
|
|
||||||
}, null, 8, ["header"]),
|
|
||||||
createVNode(unref(script), { pt: {
|
|
||||||
bodyCell: "flex items-center justify-end"
|
|
||||||
} }, {
|
|
||||||
body: withCtx((slotProps) => [
|
|
||||||
createVNode(unref(script$1), {
|
|
||||||
modelValue: editingEnabledExtensions.value[slotProps.data.name],
|
|
||||||
"onUpdate:modelValue": /* @__PURE__ */ __name(($event) => editingEnabledExtensions.value[slotProps.data.name] = $event, "onUpdate:modelValue"),
|
|
||||||
onChange: updateExtensionStatus
|
|
||||||
}, null, 8, ["modelValue", "onUpdate:modelValue"])
|
|
||||||
]),
|
|
||||||
_: 1
|
|
||||||
})
|
|
||||||
]),
|
|
||||||
_: 1
|
|
||||||
}, 8, ["value"]),
|
|
||||||
createBaseVNode("div", _hoisted_2, [
|
|
||||||
hasChanges.value ? (openBlock(), createBlock(unref(script$3), {
|
|
||||||
key: 0,
|
|
||||||
severity: "info"
|
|
||||||
}, {
|
|
||||||
default: withCtx(() => [
|
|
||||||
createBaseVNode("ul", null, [
|
|
||||||
(openBlock(true), createElementBlock(Fragment, null, renderList(changedExtensions.value, (ext) => {
|
|
||||||
return openBlock(), createElementBlock("li", {
|
|
||||||
key: ext.name
|
|
||||||
}, [
|
|
||||||
createBaseVNode("span", null, toDisplayString(unref(extensionStore).isExtensionEnabled(ext.name) ? "[-]" : "[+]"), 1),
|
|
||||||
createTextVNode(" " + toDisplayString(ext.name), 1)
|
|
||||||
]);
|
|
||||||
}), 128))
|
|
||||||
])
|
|
||||||
]),
|
|
||||||
_: 1
|
|
||||||
})) : createCommentVNode("", true),
|
|
||||||
createVNode(unref(script$4), {
|
|
||||||
label: _ctx.$t("reloadToApplyChanges"),
|
|
||||||
icon: "pi pi-refresh",
|
|
||||||
onClick: applyChanges,
|
|
||||||
disabled: !hasChanges.value,
|
|
||||||
text: "",
|
|
||||||
fluid: "",
|
|
||||||
severity: "danger"
|
|
||||||
}, null, 8, ["label", "disabled"])
|
|
||||||
])
|
|
||||||
]);
|
|
||||||
};
|
|
||||||
}
|
|
||||||
});
|
|
||||||
export {
|
|
||||||
_sfc_main as default
|
|
||||||
};
|
|
||||||
//# sourceMappingURL=ExtensionPanel-CfMfcLgI.js.map
|
|
||||||
1
web/assets/ExtensionPanel-CfMfcLgI.js.map
generated
vendored
1
web/assets/ExtensionPanel-CfMfcLgI.js.map
generated
vendored
@ -1 +0,0 @@
|
|||||||
{"version":3,"file":"ExtensionPanel-CfMfcLgI.js","sources":["../../src/components/dialog/content/setting/ExtensionPanel.vue"],"sourcesContent":["<template>\n <div class=\"extension-panel\">\n <DataTable :value=\"extensionStore.extensions\" stripedRows size=\"small\">\n <Column field=\"name\" :header=\"$t('extensionName')\" sortable></Column>\n <Column\n :pt=\"{\n bodyCell: 'flex items-center justify-end'\n }\"\n >\n <template #body=\"slotProps\">\n <ToggleSwitch\n v-model=\"editingEnabledExtensions[slotProps.data.name]\"\n @change=\"updateExtensionStatus\"\n />\n </template>\n </Column>\n </DataTable>\n <div class=\"mt-4\">\n <Message v-if=\"hasChanges\" severity=\"info\">\n <ul>\n <li v-for=\"ext in changedExtensions\" :key=\"ext.name\">\n <span>\n {{ extensionStore.isExtensionEnabled(ext.name) ? '[-]' : '[+]' }}\n </span>\n {{ ext.name }}\n </li>\n </ul>\n </Message>\n <Button\n :label=\"$t('reloadToApplyChanges')\"\n icon=\"pi pi-refresh\"\n @click=\"applyChanges\"\n :disabled=\"!hasChanges\"\n text\n fluid\n severity=\"danger\"\n />\n </div>\n </div>\n</template>\n\n<script setup lang=\"ts\">\nimport { ref, computed, onMounted } from 'vue'\nimport { useExtensionStore } from '@/stores/extensionStore'\nimport { useSettingStore } from '@/stores/settingStore'\nimport DataTable from 'primevue/datatable'\nimport Column from 'primevue/column'\nimport ToggleSwitch from 'primevue/toggleswitch'\nimport Button from 'primevue/button'\nimport Message from 'primevue/message'\n\nconst extensionStore = useExtensionStore()\nconst settingStore = useSettingStore()\n\nconst editingEnabledExtensions = ref<Record<string, boolean>>({})\n\nonMounted(() => {\n extensionStore.extensions.forEach((ext) => {\n editingEnabledExtensions.value[ext.name] =\n extensionStore.isExtensionEnabled(ext.name)\n })\n})\n\nconst changedExtensions = computed(() => {\n return extensionStore.extensions.filter(\n (ext) =>\n editingEnabledExtensions.value[ext.name] !==\n extensionStore.isExtensionEnabled(ext.name)\n )\n})\n\nconst hasChanges = computed(() => {\n return changedExtensions.value.length > 0\n})\n\nconst updateExtensionStatus = () => {\n const editingDisabledExtensionNames = Object.entries(\n editingEnabledExtensions.value\n )\n .filter(([_, enabled]) => !enabled)\n .map(([name]) => name)\n\n settingStore.set('Comfy.Extension.Disabled', [\n ...extensionStore.inactiveDisabledExtensionNames,\n ...editingDisabledExtensionNames\n ])\n}\n\nconst applyChanges = () => {\n // Refresh the page to apply changes\n window.location.reload()\n}\n</script>\n"],"names":[],"mappings":";;;;;;;;;;AAmDA,UAAM,iBAAiB;AACvB,UAAM,eAAe;AAEf,UAAA,2BAA2B,IAA6B,CAAA,CAAE;AAEhE,cAAU,MAAM;AACC,qBAAA,WAAW,QAAQ,CAAC,QAAQ;AACzC,iCAAyB,MAAM,IAAI,IAAI,IACrC,eAAe,mBAAmB,IAAI,IAAI;AAAA,MAAA,CAC7C;AAAA,IAAA,CACF;AAEK,UAAA,oBAAoB,SAAS,MAAM;AACvC,aAAO,eAAe,WAAW;AAAA,QAC/B,CAAC,QACC,yBAAyB,MAAM,IAAI,IAAI,MACvC,eAAe,mBAAmB,IAAI,IAAI;AAAA,MAAA;AAAA,IAC9C,CACD;AAEK,UAAA,aAAa,SAAS,MAAM;AACzB,aAAA,kBAAkB,MAAM,SAAS;AAAA,IAAA,CACzC;AAED,UAAM,wBAAwB,6BAAM;AAClC,YAAM,gCAAgC,OAAO;AAAA,QAC3C,yBAAyB;AAAA,MAExB,EAAA,OAAO,CAAC,CAAC,GAAG,OAAO,MAAM,CAAC,OAAO,EACjC,IAAI,CAAC,CAAC,IAAI,MAAM,IAAI;AAEvB,mBAAa,IAAI,4BAA4B;AAAA,QAC3C,GAAG,eAAe;AAAA,QAClB,GAAG;AAAA,MAAA,CACJ;AAAA,IAAA,GAV2B;AAa9B,UAAM,eAAe,6BAAM;AAEzB,aAAO,SAAS;IAAO,GAFJ;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;"}
|
|
||||||
117
web/assets/ExtensionPanel-DsD42OtO.js
generated
vendored
Normal file
117
web/assets/ExtensionPanel-DsD42OtO.js
generated
vendored
Normal file
@ -0,0 +1,117 @@
|
|||||||
|
var __defProp = Object.defineProperty;
|
||||||
|
var __name = (target, value) => __defProp(target, "name", { value, configurable: true });
|
||||||
|
import { d as defineComponent, r as ref, c6 as FilterMatchMode, ca as useExtensionStore, u as useSettingStore, o as onMounted, q as computed, g as openBlock, x as createBlock, y as withCtx, i as createVNode, c7 as SearchBox, z as unref, bT as script, A as createBaseVNode, h as createElementBlock, O as renderList, a6 as toDisplayString, aw as createTextVNode, N as Fragment, D as script$1, j as createCommentVNode, bV as script$3, c8 as _sfc_main$1 } from "./index-CoOvI8ZH.js";
|
||||||
|
import { s as script$2, a as script$4 } from "./index-DK6Kev7f.js";
|
||||||
|
import "./index-D4DWQPPQ.js";
|
||||||
|
const _hoisted_1 = { class: "flex justify-end" };
|
||||||
|
const _sfc_main = /* @__PURE__ */ defineComponent({
|
||||||
|
__name: "ExtensionPanel",
|
||||||
|
setup(__props) {
|
||||||
|
const filters = ref({
|
||||||
|
global: { value: "", matchMode: FilterMatchMode.CONTAINS }
|
||||||
|
});
|
||||||
|
const extensionStore = useExtensionStore();
|
||||||
|
const settingStore = useSettingStore();
|
||||||
|
const editingEnabledExtensions = ref({});
|
||||||
|
onMounted(() => {
|
||||||
|
extensionStore.extensions.forEach((ext) => {
|
||||||
|
editingEnabledExtensions.value[ext.name] = extensionStore.isExtensionEnabled(ext.name);
|
||||||
|
});
|
||||||
|
});
|
||||||
|
const changedExtensions = computed(() => {
|
||||||
|
return extensionStore.extensions.filter(
|
||||||
|
(ext) => editingEnabledExtensions.value[ext.name] !== extensionStore.isExtensionEnabled(ext.name)
|
||||||
|
);
|
||||||
|
});
|
||||||
|
const hasChanges = computed(() => {
|
||||||
|
return changedExtensions.value.length > 0;
|
||||||
|
});
|
||||||
|
const updateExtensionStatus = /* @__PURE__ */ __name(() => {
|
||||||
|
const editingDisabledExtensionNames = Object.entries(
|
||||||
|
editingEnabledExtensions.value
|
||||||
|
).filter(([_, enabled]) => !enabled).map(([name]) => name);
|
||||||
|
settingStore.set("Comfy.Extension.Disabled", [
|
||||||
|
...extensionStore.inactiveDisabledExtensionNames,
|
||||||
|
...editingDisabledExtensionNames
|
||||||
|
]);
|
||||||
|
}, "updateExtensionStatus");
|
||||||
|
const applyChanges = /* @__PURE__ */ __name(() => {
|
||||||
|
window.location.reload();
|
||||||
|
}, "applyChanges");
|
||||||
|
return (_ctx, _cache) => {
|
||||||
|
return openBlock(), createBlock(_sfc_main$1, {
|
||||||
|
value: "Extension",
|
||||||
|
class: "extension-panel"
|
||||||
|
}, {
|
||||||
|
header: withCtx(() => [
|
||||||
|
createVNode(SearchBox, {
|
||||||
|
modelValue: filters.value["global"].value,
|
||||||
|
"onUpdate:modelValue": _cache[0] || (_cache[0] = ($event) => filters.value["global"].value = $event),
|
||||||
|
placeholder: _ctx.$t("searchExtensions") + "..."
|
||||||
|
}, null, 8, ["modelValue", "placeholder"]),
|
||||||
|
hasChanges.value ? (openBlock(), createBlock(unref(script), {
|
||||||
|
key: 0,
|
||||||
|
severity: "info",
|
||||||
|
"pt:text": "w-full"
|
||||||
|
}, {
|
||||||
|
default: withCtx(() => [
|
||||||
|
createBaseVNode("ul", null, [
|
||||||
|
(openBlock(true), createElementBlock(Fragment, null, renderList(changedExtensions.value, (ext) => {
|
||||||
|
return openBlock(), createElementBlock("li", {
|
||||||
|
key: ext.name
|
||||||
|
}, [
|
||||||
|
createBaseVNode("span", null, toDisplayString(unref(extensionStore).isExtensionEnabled(ext.name) ? "[-]" : "[+]"), 1),
|
||||||
|
createTextVNode(" " + toDisplayString(ext.name), 1)
|
||||||
|
]);
|
||||||
|
}), 128))
|
||||||
|
]),
|
||||||
|
createBaseVNode("div", _hoisted_1, [
|
||||||
|
createVNode(unref(script$1), {
|
||||||
|
label: _ctx.$t("reloadToApplyChanges"),
|
||||||
|
onClick: applyChanges,
|
||||||
|
outlined: "",
|
||||||
|
severity: "danger"
|
||||||
|
}, null, 8, ["label"])
|
||||||
|
])
|
||||||
|
]),
|
||||||
|
_: 1
|
||||||
|
})) : createCommentVNode("", true)
|
||||||
|
]),
|
||||||
|
default: withCtx(() => [
|
||||||
|
createVNode(unref(script$4), {
|
||||||
|
value: unref(extensionStore).extensions,
|
||||||
|
stripedRows: "",
|
||||||
|
size: "small",
|
||||||
|
filters: filters.value
|
||||||
|
}, {
|
||||||
|
default: withCtx(() => [
|
||||||
|
createVNode(unref(script$2), {
|
||||||
|
field: "name",
|
||||||
|
header: _ctx.$t("extensionName"),
|
||||||
|
sortable: ""
|
||||||
|
}, null, 8, ["header"]),
|
||||||
|
createVNode(unref(script$2), { pt: {
|
||||||
|
bodyCell: "flex items-center justify-end"
|
||||||
|
} }, {
|
||||||
|
body: withCtx((slotProps) => [
|
||||||
|
createVNode(unref(script$3), {
|
||||||
|
modelValue: editingEnabledExtensions.value[slotProps.data.name],
|
||||||
|
"onUpdate:modelValue": /* @__PURE__ */ __name(($event) => editingEnabledExtensions.value[slotProps.data.name] = $event, "onUpdate:modelValue"),
|
||||||
|
onChange: updateExtensionStatus
|
||||||
|
}, null, 8, ["modelValue", "onUpdate:modelValue"])
|
||||||
|
]),
|
||||||
|
_: 1
|
||||||
|
})
|
||||||
|
]),
|
||||||
|
_: 1
|
||||||
|
}, 8, ["value", "filters"])
|
||||||
|
]),
|
||||||
|
_: 1
|
||||||
|
});
|
||||||
|
};
|
||||||
|
}
|
||||||
|
});
|
||||||
|
export {
|
||||||
|
_sfc_main as default
|
||||||
|
};
|
||||||
|
//# sourceMappingURL=ExtensionPanel-DsD42OtO.js.map
|
||||||
1
web/assets/ExtensionPanel-DsD42OtO.js.map
generated
vendored
Normal file
1
web/assets/ExtensionPanel-DsD42OtO.js.map
generated
vendored
Normal file
@ -0,0 +1 @@
|
|||||||
|
{"version":3,"file":"ExtensionPanel-DsD42OtO.js","sources":["../../src/components/dialog/content/setting/ExtensionPanel.vue"],"sourcesContent":["<template>\n <PanelTemplate value=\"Extension\" class=\"extension-panel\">\n <template #header>\n <SearchBox\n v-model=\"filters['global'].value\"\n :placeholder=\"$t('searchExtensions') + '...'\"\n />\n <Message v-if=\"hasChanges\" severity=\"info\" pt:text=\"w-full\">\n <ul>\n <li v-for=\"ext in changedExtensions\" :key=\"ext.name\">\n <span>\n {{ extensionStore.isExtensionEnabled(ext.name) ? '[-]' : '[+]' }}\n </span>\n {{ ext.name }}\n </li>\n </ul>\n <div class=\"flex justify-end\">\n <Button\n :label=\"$t('reloadToApplyChanges')\"\n @click=\"applyChanges\"\n outlined\n severity=\"danger\"\n />\n </div>\n </Message>\n </template>\n <DataTable\n :value=\"extensionStore.extensions\"\n stripedRows\n size=\"small\"\n :filters=\"filters\"\n >\n <Column field=\"name\" :header=\"$t('extensionName')\" sortable></Column>\n <Column\n :pt=\"{\n bodyCell: 'flex items-center justify-end'\n }\"\n >\n <template #body=\"slotProps\">\n <ToggleSwitch\n v-model=\"editingEnabledExtensions[slotProps.data.name]\"\n @change=\"updateExtensionStatus\"\n />\n </template>\n </Column>\n </DataTable>\n </PanelTemplate>\n</template>\n\n<script setup lang=\"ts\">\nimport { ref, computed, onMounted } from 'vue'\nimport { useExtensionStore } from '@/stores/extensionStore'\nimport { useSettingStore } from '@/stores/settingStore'\nimport DataTable from 'primevue/datatable'\nimport Column from 'primevue/column'\nimport ToggleSwitch from 'primevue/toggleswitch'\nimport Button from 'primevue/button'\nimport Message from 'primevue/message'\nimport { FilterMatchMode } from '@primevue/core/api'\nimport PanelTemplate from './PanelTemplate.vue'\nimport SearchBox from '@/components/common/SearchBox.vue'\n\nconst filters = ref({\n global: { value: '', matchMode: FilterMatchMode.CONTAINS }\n})\n\nconst extensionStore = useExtensionStore()\nconst settingStore = useSettingStore()\n\nconst editingEnabledExtensions = ref<Record<string, boolean>>({})\n\nonMounted(() => {\n extensionStore.extensions.forEach((ext) => {\n editingEnabledExtensions.value[ext.name] =\n extensionStore.isExtensionEnabled(ext.name)\n })\n})\n\nconst changedExtensions = computed(() => {\n return extensionStore.extensions.filter(\n (ext) =>\n editingEnabledExtensions.value[ext.name] !==\n extensionStore.isExtensionEnabled(ext.name)\n )\n})\n\nconst hasChanges = computed(() => {\n return changedExtensions.value.length > 0\n})\n\nconst updateExtensionStatus = () => {\n const editingDisabledExtensionNames = Object.entries(\n editingEnabledExtensions.value\n )\n .filter(([_, enabled]) => !enabled)\n .map(([name]) => name)\n\n settingStore.set('Comfy.Extension.Disabled', [\n ...extensionStore.inactiveDisabledExtensionNames,\n ...editingDisabledExtensionNames\n ])\n}\n\nconst applyChanges = () => {\n // Refresh the page to apply changes\n window.location.reload()\n}\n</script>\n"],"names":[],"mappings":";;;;;;;;;AA8DA,UAAM,UAAU,IAAI;AAAA,MAClB,QAAQ,EAAE,OAAO,IAAI,WAAW,gBAAgB,SAAS;AAAA,IAAA,CAC1D;AAED,UAAM,iBAAiB;AACvB,UAAM,eAAe;AAEf,UAAA,2BAA2B,IAA6B,CAAA,CAAE;AAEhE,cAAU,MAAM;AACC,qBAAA,WAAW,QAAQ,CAAC,QAAQ;AACzC,iCAAyB,MAAM,IAAI,IAAI,IACrC,eAAe,mBAAmB,IAAI,IAAI;AAAA,MAAA,CAC7C;AAAA,IAAA,CACF;AAEK,UAAA,oBAAoB,SAAS,MAAM;AACvC,aAAO,eAAe,WAAW;AAAA,QAC/B,CAAC,QACC,yBAAyB,MAAM,IAAI,IAAI,MACvC,eAAe,mBAAmB,IAAI,IAAI;AAAA,MAAA;AAAA,IAC9C,CACD;AAEK,UAAA,aAAa,SAAS,MAAM;AACzB,aAAA,kBAAkB,MAAM,SAAS;AAAA,IAAA,CACzC;AAED,UAAM,wBAAwB,6BAAM;AAClC,YAAM,gCAAgC,OAAO;AAAA,QAC3C,yBAAyB;AAAA,MAExB,EAAA,OAAO,CAAC,CAAC,GAAG,OAAO,MAAM,CAAC,OAAO,EACjC,IAAI,CAAC,CAAC,IAAI,MAAM,IAAI;AAEvB,mBAAa,IAAI,4BAA4B;AAAA,QAC3C,GAAG,eAAe;AAAA,QAClB,GAAG;AAAA,MAAA,CACJ;AAAA,IAAA,GAV2B;AAa9B,UAAM,eAAe,6BAAM;AAEzB,aAAO,SAAS;IAAO,GAFJ;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;;"}
|
||||||
1
web/assets/GraphView-BCOd0Zle.js.map
generated
vendored
1
web/assets/GraphView-BCOd0Zle.js.map
generated
vendored
File diff suppressed because one or more lines are too long
572
web/assets/GraphView-BCOd0Zle.js → web/assets/GraphView-BW5soyxY.js
generated
vendored
572
web/assets/GraphView-BCOd0Zle.js → web/assets/GraphView-BW5soyxY.js
generated
vendored
File diff suppressed because one or more lines are too long
1
web/assets/GraphView-BW5soyxY.js.map
generated
vendored
Normal file
1
web/assets/GraphView-BW5soyxY.js.map
generated
vendored
Normal file
File diff suppressed because one or more lines are too long
50
web/assets/GraphView-CghYAxkP.css → web/assets/GraphView-DtkYXy38.css
generated
vendored
50
web/assets/GraphView-CghYAxkP.css → web/assets/GraphView-DtkYXy38.css
generated
vendored
@ -106,32 +106,6 @@
|
|||||||
margin: -0.125rem 0.125rem;
|
margin: -0.125rem 0.125rem;
|
||||||
}
|
}
|
||||||
|
|
||||||
.comfy-vue-node-search-container[data-v-2d409367] {
|
|
||||||
display: flex;
|
|
||||||
width: 100%;
|
|
||||||
min-width: 26rem;
|
|
||||||
align-items: center;
|
|
||||||
justify-content: center;
|
|
||||||
}
|
|
||||||
.comfy-vue-node-search-container[data-v-2d409367] * {
|
|
||||||
pointer-events: auto;
|
|
||||||
}
|
|
||||||
.comfy-vue-node-preview-container[data-v-2d409367] {
|
|
||||||
position: absolute;
|
|
||||||
left: -350px;
|
|
||||||
top: 50px;
|
|
||||||
}
|
|
||||||
.comfy-vue-node-search-box[data-v-2d409367] {
|
|
||||||
z-index: 10;
|
|
||||||
flex-grow: 1;
|
|
||||||
}
|
|
||||||
._filter-button[data-v-2d409367] {
|
|
||||||
z-index: 10;
|
|
||||||
}
|
|
||||||
._dialog[data-v-2d409367] {
|
|
||||||
min-width: 26rem;
|
|
||||||
}
|
|
||||||
|
|
||||||
.invisible-dialog-root {
|
.invisible-dialog-root {
|
||||||
width: 60%;
|
width: 60%;
|
||||||
min-width: 24rem;
|
min-width: 24rem;
|
||||||
@ -184,10 +158,10 @@
|
|||||||
z-index: 9999;
|
z-index: 9999;
|
||||||
}
|
}
|
||||||
|
|
||||||
[data-v-9eb975c3] .p-togglebutton::before {
|
[data-v-783f8efe] .p-togglebutton::before {
|
||||||
display: none
|
display: none
|
||||||
}
|
}
|
||||||
[data-v-9eb975c3] .p-togglebutton {
|
[data-v-783f8efe] .p-togglebutton {
|
||||||
position: relative;
|
position: relative;
|
||||||
flex-shrink: 0;
|
flex-shrink: 0;
|
||||||
border-radius: 0px;
|
border-radius: 0px;
|
||||||
@ -195,14 +169,14 @@
|
|||||||
padding-left: 0.5rem;
|
padding-left: 0.5rem;
|
||||||
padding-right: 0.5rem
|
padding-right: 0.5rem
|
||||||
}
|
}
|
||||||
[data-v-9eb975c3] .p-togglebutton.p-togglebutton-checked {
|
[data-v-783f8efe] .p-togglebutton.p-togglebutton-checked {
|
||||||
border-bottom-width: 2px;
|
border-bottom-width: 2px;
|
||||||
border-bottom-color: var(--p-button-text-primary-color)
|
border-bottom-color: var(--p-button-text-primary-color)
|
||||||
}
|
}
|
||||||
[data-v-9eb975c3] .p-togglebutton-checked .close-button,[data-v-9eb975c3] .p-togglebutton:hover .close-button {
|
[data-v-783f8efe] .p-togglebutton-checked .close-button,[data-v-783f8efe] .p-togglebutton:hover .close-button {
|
||||||
visibility: visible
|
visibility: visible
|
||||||
}
|
}
|
||||||
.status-indicator[data-v-9eb975c3] {
|
.status-indicator[data-v-783f8efe] {
|
||||||
position: absolute;
|
position: absolute;
|
||||||
font-weight: 700;
|
font-weight: 700;
|
||||||
font-size: 1.5rem;
|
font-size: 1.5rem;
|
||||||
@ -210,10 +184,10 @@
|
|||||||
left: 50%;
|
left: 50%;
|
||||||
transform: translate(-50%, -50%)
|
transform: translate(-50%, -50%)
|
||||||
}
|
}
|
||||||
[data-v-9eb975c3] .p-togglebutton:hover .status-indicator {
|
[data-v-783f8efe] .p-togglebutton:hover .status-indicator {
|
||||||
display: none
|
display: none
|
||||||
}
|
}
|
||||||
[data-v-9eb975c3] .p-togglebutton .close-button {
|
[data-v-783f8efe] .p-togglebutton .close-button {
|
||||||
visibility: hidden
|
visibility: hidden
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -241,26 +215,26 @@
|
|||||||
border-bottom-right-radius: 0;
|
border-bottom-right-radius: 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
.actionbar[data-v-eb6e9acf] {
|
.actionbar[data-v-542a7001] {
|
||||||
pointer-events: all;
|
pointer-events: all;
|
||||||
position: fixed;
|
position: fixed;
|
||||||
z-index: 1000;
|
z-index: 1000;
|
||||||
}
|
}
|
||||||
.actionbar.is-docked[data-v-eb6e9acf] {
|
.actionbar.is-docked[data-v-542a7001] {
|
||||||
position: static;
|
position: static;
|
||||||
border-style: none;
|
border-style: none;
|
||||||
background-color: transparent;
|
background-color: transparent;
|
||||||
padding: 0px;
|
padding: 0px;
|
||||||
}
|
}
|
||||||
.actionbar.is-dragging[data-v-eb6e9acf] {
|
.actionbar.is-dragging[data-v-542a7001] {
|
||||||
-webkit-user-select: none;
|
-webkit-user-select: none;
|
||||||
-moz-user-select: none;
|
-moz-user-select: none;
|
||||||
user-select: none;
|
user-select: none;
|
||||||
}
|
}
|
||||||
[data-v-eb6e9acf] .p-panel-content {
|
[data-v-542a7001] .p-panel-content {
|
||||||
padding: 0.25rem;
|
padding: 0.25rem;
|
||||||
}
|
}
|
||||||
[data-v-eb6e9acf] .p-panel-header {
|
[data-v-542a7001] .p-panel-header {
|
||||||
display: none;
|
display: none;
|
||||||
}
|
}
|
||||||
|
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Loading…
x
Reference in New Issue
Block a user