27 Commits
Author SHA1 Message Date
Robin Huangandsnomiao ab4a4fb38c chore(publish): update GitHub Actions workflow for node publishing (#31)
- Add permissions for issue writing
- Set condition to run job only for 'Jannchie' repository owner
- Update action version from 'main' to 'v1' for stability and consistency

Co-authored-by: snomiao <snomiao+comfy-pr@gmail.com>
2025-04-07 18:03:24 +09:00
Jianqi Pan 92b0b83839 fix(average): get average color 2024-09-15 00:22:25 +09:00
Jianqi Pan 17a0274d77 version: v1.1.0 2024-07-31 22:12:38 +09:00
Jianqi Pan 3c121b126c fix: enhance the compatibility between diffusers and comfy 2024-07-31 22:12:08 +09:00
haohaocreates 4f6bb64067 Add Github Action for Publishing to Comfy Registry (#19) 2024-06-20 21:08:09 +09:00
Jianqi Pan 4b2aac5426 Merge pull request #20 from haohaocreates/pyproject
Add pyproject.toml for Custom Node Registry
2024-06-20 21:07:52 +09:00
haohaocreates 59326f53d1 chore(pyproject): Add pyproject.toml for Custom Node Registry 2024-05-22 20:16:07 -04:00
Jianqi Pan aa0211ee28 🔧 chore: relaese vmem before update pipeline 2024-05-08 20:12:10 +09:00
Jianqi Pan 55636490a3 🔧 chore: relaese vmem after generate 2024-05-08 20:10:17 +09:00
Jianqi Pan 8698c62af7 🩹 fix(pipeline): no ip adapter 2024-05-01 16:26:02 +09:00
Jianqi Pan d3f732236a ✨ feat(tiny-vae): support tiny vae 2024-04-30 20:56:45 +09:00
Jianqi Pan 5f94a0cf46 ✨ feat(tiny-vae): support tiny vae 2024-04-30 20:56:30 +09:00
Jianqi Pan 371e006976 ✨ feat(ip-adapter): support ip adapter 2024-04-30 20:55:58 +09:00
Jianqi Pan 49c360978a 🩹 fix(type): fix type def 2024-04-22 02:21:47 +09:00
Jianqi Pan f3c696d5db ✨ feat(example): add QR code example 2024-04-18 01:25:51 +09:00
Jianqi Pan d3edc4a8f4 🩹 fix(dtype): collect unet & vae type 2024-03-24 15:35:14 +09:00
Jianqi Pan 5be339b5a0 📚 docs: update readme 2024-03-23 23:16:37 +09:00
Jianqi Pan 0d9da4ebc6 📚 docs: update examples 2024-03-23 21:24:24 +09:00
Jianqi Pan 83c9d84e3e 🔨 refactor(controlnet): better controlnet node to auto download
model
2024-03-23 21:14:05 +09:00
Jianqi Pan 83263b4f96 🔨 refactor(ref): better reference only 2024-03-23 21:13:41 +09:00
Jianqi Pan 4b485552be Merge branch 'main' of github.com:Jannchie/ComfyUI-J 2024-03-23 16:47:08 +09:00
Jianqi Pan e939616807 📚 docs: update todo list & add desc 2024-03-23 16:46:49 +09:00
Jianqi Pan ad468be7ac Rename Inpainting.png to inpainting.png 2024-03-23 16:42:50 +09:00
Jianqi Pan a211e1c623 🩹 fix(mask): do not resize mask 2024-03-23 16:39:08 +09:00
Jianqi Pan bb431b9616 📚 docs: fix typo & file name errors 2024-03-23 16:10:57 +09:00
Jianqi Pan 0574d49da5 ✨ feat(faq): add faq 2024-03-22 21:41:19 +09:00
Jianqi Pan d0f73cc800 📚 docs: add more examples 2024-03-22 21:29:36 +09:00
14 changed files with 450 additions and 602 deletions
+25
View File
@@ -0,0 +1,25 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
permissions:
issues: write
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Jannchie' }}
steps:
- name: Check out code
uses: actions/checkout@v4
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@v1
with:
## Add your own personal access token to your Github Repository secrets and reference it here.
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+76
View File
@@ -1,6 +1,82 @@
# ComfyUI-J # ComfyUI-J
## Introduction
Jannchie's ComfyUI custom nodes. Jannchie's ComfyUI custom nodes.
This is a completely different set of nodes than Comfy's own KSampler series. This is a completely different set of nodes than Comfy's own KSampler series.
This set of nodes is based on Diffusers, which makes it easier to import models, apply prompts with weights, inpaint, reference only, controlnet, etc. This set of nodes is based on Diffusers, which makes it easier to import models, apply prompts with weights, inpaint, reference only, controlnet, etc.
## Installation
In the `custom_nodes` directory, run
```bash
git clone https://github.com/Jannchie/ComfyUI-J
cd ComfyUI-J
pip install -r requirements.txt
```
## Examples
### Base Usage of Jannchie's Diffusers Pipeline
You only have to deal with 4 nodes. The default comfy workflow uses 7 nodes to achieve the same result.
![Base Usage](./examples/base.png)
### Reference Only with Jannchie's Diffusers Pipeline
ref_only supports two modes: attn and attn + adain, and can adjust the style fidelity parameter to control the style.
![Reference only](./examples/reference_only.png)
### ControlNet with Jannchie's Diffusers Pipeline
ContorlNet is also easier to use. A DiffusersControlnetLoader node is provided for loading models. This node automatically detects if the corresponding ControlNet has been downloaded locally, and pulls the model from the huggingface if it has not.
![ControlNet](./examples/controlnet.png)
## Inpainting with Jannchie's Diffusers Pipeline
![Inpainting](./examples/inpainting.png)
## Remove something with Jannchie's Diffusers Pipeline
![Remove something](./examples/remove_something.png)
## Change Clothes with Jannchie's Diffusers Pipeline
This is a composite application of diffusers pipeline custom node. Includes:
- Reference only
- ControlNet
- Inpainting
- Textual Inversion
This is a demonstration of a simple workflow for properly dressing a character.
A checkpoint for stablediffusion 1.5 is all your need. But for full automation, I use the `Comfyui_segformer_b2_clothes` custom node for generating masks. you can draw your own masks without it.
![Change Clothes](./examples/change_clothes.png)
## QR Code
![QR Code](./examples/qr_code.png)
## FAQ
### Why Diffusers?
Unlike Web UI and Comfy, Diffusers is an image generation tool for researchers. It has a large ecosystem, a clearer code structure and a simpler interface.
ComfyUI's KSampler is nice, but some of the features are incomplete or hard to be access, it's 2042 and I still haven't found a good Reference Only implementation; Inpaint also works differently than I thought it would; I don't understand at all why ControlNet's nodes need to pass in a CLIP; and I don't want to deal with what's going on with Latent, please just return an Image instead of making me decode it with a vae. Diffusers provides a pipeline wrapper that makes generation a lot easier.
### Why ComfyUI?
But combining research results is not an easy task, Comfy is good at combining and sharing combinations with others. While debugging custom nodes as a developer can be a pain, using Comfy makes it faster to verify and share.
## TODO
- [ ] Add LoRA support
- [ ] Stable Diffusion XL support
+195 -49
View File
@@ -1,4 +1,5 @@
import contextlib import contextlib
import gc
import random import random
from collections import Counter from collections import Counter
@@ -74,6 +75,9 @@ def latents_to_img_tensor(pipeline, latents):
# 1. 输入的 latents 是一个 -1 ~ 1 之间的 tensor # 1. 输入的 latents 是一个 -1 ~ 1 之间的 tensor
# 2. 先进行缩放 # 2. 先进行缩放
scaled_latents = latents / pipeline.vae.config.scaling_factor scaled_latents = latents / pipeline.vae.config.scaling_factor
# 转成 vae 类型
scaled_latents = scaled_latents.to(dtype=comfy.model_management.vae_dtype())
print(scaled_latents.dtype, pipeline.vae.dtype)
# 3. 解码,返回的是 -1 ~ 1 之间的 tensor # 3. 解码,返回的是 -1 ~ 1 之间的 tensor
dec_tensor = pipeline.vae.decode(scaled_latents, return_dict=False)[0] dec_tensor = pipeline.vae.decode(scaled_latents, return_dict=False)[0]
# 4. 缩放到 0 ~ 1 之间 # 4. 缩放到 0 ~ 1 之间
@@ -197,7 +201,7 @@ class GetFilledColorImage:
"default": 0.0, "default": 0.0,
"min": 0.0, "min": 0.0,
"max": 1.0, "max": 1.0,
"step": 0.1, "step": 0.01,
"display": "number", "display": "number",
}, },
), ),
@@ -207,7 +211,7 @@ class GetFilledColorImage:
"default": 0.0, "default": 0.0,
"min": 0.0, "min": 0.0,
"max": 1.0, "max": 1.0,
"step": 0.1, "step": 0.01,
"display": "number", "display": "number",
}, },
), ),
@@ -217,7 +221,7 @@ class GetFilledColorImage:
"default": 0.0, "default": 0.0,
"min": 0.0, "min": 0.0,
"max": 1.0, "max": 1.0,
"step": 0.1, "step": 0.01,
"display": "number", "display": "number",
}, },
), ),
@@ -283,6 +287,7 @@ class DiffusersTextureInversionLoader:
path = folder_paths.get_full_path("embeddings", texture_inversion) path = folder_paths.get_full_path("embeddings", texture_inversion)
token = texture_inversion.split(".")[0] token = texture_inversion.split(".")[0]
pipeline.load_textual_inversion(path, token=token) pipeline.load_textual_inversion(path, token=token)
print(f"Loaded {texture_inversion}")
return (pipeline,) return (pipeline,)
@@ -297,7 +302,7 @@ class GetAverageColorFromImage:
return { return {
"required": { "required": {
"image": ("IMAGE",), "image": ("IMAGE",),
"average": ("STRING", {"default": "mean", "options": ["mean", "mode"]}), "average": (("mean", "mode"),),
}, },
"optional": { "optional": {
"mask": ("MASK",), "mask": ("MASK",),
@@ -305,48 +310,112 @@ class GetAverageColorFromImage:
} }
def run(self, image: torch.Tensor, average: str, mask: torch.Tensor = None): def run(self, image: torch.Tensor, average: str, mask: torch.Tensor = None):
if mask is not None:
assert (
mask.ndim == image.ndim - 1
), "Mask dimensions must be one less than image dimensions."
mask = mask.unsqueeze(3) # Unsqueeze to match (B, 1, H, W)
if mask is not None and torch.sum(mask) == 0:
mask = None
if average == "mean": if average == "mean":
return self.run_avg(image, mask) return self.run_avg(image, mask)
elif average == "mode": elif average == "mode":
return self.run_mode(image, mask) return self.run_mode(image, mask)
else:
raise ValueError("average must be either 'mean' or 'mode'")
def run_avg(self, image: torch.Tensor, mask: torch.Tensor = None): def run_avg(self, image: torch.Tensor, mask: torch.Tensor = None):
if mask is not None:
mask = mask.unsqueeze(1)
masked_image = image * mask if mask is not None else image masked_image = image * mask if mask is not None else image
pixel_sum = torch.sum(masked_image, dim=(2, 3))
pixel_count = (
torch.sum(mask, dim=(2, 3))
if mask is not None
else torch.prod(torch.tensor(image.shape[2:]))
)
average_rgb = pixel_sum / pixel_count.unsqueeze(1)
average_rgb = torch.round(average_rgb) pixel_sum = torch.sum(masked_image, dim=(1, 2))
if mask is not None:
return tuple(average_rgb.squeeze().tolist()) pixel_count = torch.sum(mask, dim=(1, 2)).unsqueeze(1)
else:
pixel_count = torch.tensor(image.shape[1] * image.shape[2]).unsqueeze(0)
average_rgb = pixel_sum / pixel_count
average_rgb = torch.round(average_rgb * 255)
return tuple(average_rgb.squeeze().int().tolist())
def run_mode(self, image: torch.Tensor, mask: torch.Tensor = None): def run_mode(self, image: torch.Tensor, mask: torch.Tensor = None):
image = image.permute(0, 3, 1, 2)
if mask is not None: if mask is not None:
mask = mask.unsqueeze(1) image = image * mask
masked_image = image * mask if mask is not None else image # Flatten the image to a 2D matrix where each row is a color
pixel_values = masked_image.view( flattened_image = image.view(-1, image.shape[-1])
masked_image.shape[0], masked_image.shape[1], -1
# If mask is provided, remove rows where mask is zero
if mask is not None:
flattened_mask = mask.view(-1, 1)
flattened_image = flattened_image[flattened_mask.squeeze() > 0]
# Convert the pixel values to a format that can be efficiently counted
unique_colors, counts = torch.unique(flattened_image, return_counts=True, dim=0)
# Find the most frequent color
max_idx = torch.argmax(counts)
mode_rgb = unique_colors[max_idx]
mode_rgb = torch.round(mode_rgb * 255)
return tuple(mode_rgb.int().tolist())
class DiffusersXLPipeline:
CATEGORY = "Jannchie"
FUNCTION = "run"
RETURN_TYPES = ("DIFFUSERS_PIPELINE",)
RETURN_NAMES = ("pipeline",)
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"ckpt_name": ([],),
},
"optional": {
"vae_name": (
folder_paths.get_filename_list("vae") + ["-"],
{"default": "-"},
),
"scheduler_name": (
list(schedulers.keys()) + ["-"],
{
"default": "-",
},
),
"use_tiny_vae": (
["disable", "enable"],
{
"default": "disable",
},
),
},
}
def run(
self,
ckpt_name: str,
vae_name: str = None,
scheduler_name: str = None,
use_tiny_vae: str = "disable",
):
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
if ckpt_path is None:
ckpt_path = ckpt_name
if vae_name == "-":
vae_path = None
else:
vae_path = folder_paths.get_full_path("vae", vae_name)
if scheduler_name == "-":
scheduler_name = None
self.pipeline_wrapper = PipelineWrapper(
ckpt_path,
vae_path,
scheduler_name,
pipeline=StableDiffusionPipeline,
use_tiny_vae=use_tiny_vae == "enable",
) )
pixel_values = pixel_values.permute(0, 2, 1) return (self.pipeline_wrapper.pipeline,)
pixel_values = pixel_values.reshape(-1, pixel_values.shape[2])
pixel_values = [
tuple(color.tolist()) for color in pixel_values.numpy() if color.max() > 0
]
if not pixel_values:
return (0, 0, 0)
color_counts = Counter(pixel_values)
return max(color_counts, key=color_counts.get)
class DiffusersPipeline: class DiffusersPipeline:
@@ -372,11 +441,27 @@ class DiffusersPipeline:
"default": "-", "default": "-",
}, },
), ),
"use_tiny_vae": (
["disable", "enable"],
{
"default": "disable",
},
),
}, },
} }
def run(self, ckpt_name: str, vae_name: str = None, scheduler_name: str = None): def run(
self,
ckpt_name: str,
vae_name: str = None,
scheduler_name: str = None,
use_tiny_vae: str = "disable",
):
torch.cuda.empty_cache()
gc.collect()
ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name) ckpt_path = folder_paths.get_full_path("checkpoints", ckpt_name)
if ckpt_path is None:
ckpt_path = ckpt_name
if vae_name == "-": if vae_name == "-":
vae_path = None vae_path = None
else: else:
@@ -384,7 +469,9 @@ class DiffusersPipeline:
if scheduler_name == "-": if scheduler_name == "-":
scheduler_name = None scheduler_name = None
self.pipeline_wrapper = PipelineWrapper(ckpt_path, vae_path, scheduler_name) self.pipeline_wrapper = PipelineWrapper(
ckpt_path, vae_path, scheduler_name, use_tiny_vae=use_tiny_vae == "enable"
)
return (self.pipeline_wrapper.pipeline,) return (self.pipeline_wrapper.pipeline,)
@@ -431,7 +518,7 @@ class DiffusersPrepareLatents:
batch_size=batch_size, batch_size=batch_size,
height=height, height=height,
width=width, width=width,
dtype=comfy.model_management.VAE_DTYPE, dtype=comfy.model_management.vae_dtype(),
device=device, device=device,
generator=generator, generator=generator,
latents=latents, latents=latents,
@@ -459,6 +546,27 @@ class DiffusersDecoder:
return (res,) return (res,)
# 'https://huggingface.co/lllyasviel/ControlNet-v1-1/blob/main/control_v11p_sd15_canny.pth'
controlnet_list = [
"canny",
"openpose",
"depth",
"tile",
"ip2p",
"shuffle",
"inpaint",
"lineart",
"mlsd",
"normalbae",
"scribble",
"seg",
"softedge",
"lineart_anime",
"other",
]
class DiffusersControlNetLoader: class DiffusersControlNetLoader:
CATEGORY = "Jannchie" CATEGORY = "Jannchie"
FUNCTION = "run" FUNCTION = "run"
@@ -469,19 +577,42 @@ class DiffusersControlNetLoader:
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": { "required": {
"controlnet_model_name": ( "controlnet_model_name": (controlnet_list,),
folder_paths.get_filename_list("controlnet"), },
), "optional": {
"controlnet_model_file": (folder_paths.get_filename_list("controlnet"),)
}, },
} }
def run(self, controlnet_model_name: str): def run(self, controlnet_model_name: str, controlnet_model_file: str = ""):
controlnet_model_path = folder_paths.get_full_path( file_list = folder_paths.get_filename_list("controlnet")
"controlnet", controlnet_model_name if controlnet_model_name == "other":
) controlnet_model_path = folder_paths.get_full_path(
controlnet = ControlNetModel.from_single_file(controlnet_model_path).to( "controlnet", controlnet_model_file
)
else:
if controlnet_model_name == "depth":
file_name = f"control_v11f1p_sd15_{controlnet_model_name}.pth"
elif controlnet_model_name == "tile":
file_name = f"control_v11f1e_sd15_{controlnet_model_name}.pth"
else:
file_name = f"control_v11p_sd15_{controlnet_model_name}.pth"
controlnet_model_path = next(
(
folder_paths.get_full_path("controlnet", file)
for file in file_list
if file_name in file
),
None,
)
if controlnet_model_path is None:
controlnet_model_path = f"https://huggingface.co/lllyasviel/ControlNet-v1-1/blob/main/{file_name}"
controlnet = ControlNetModel.from_single_file(
controlnet_model_path,
cache_dir=folder_paths.get_folder_paths("controlnet")[0],
).to(
device=comfy.model_management.get_torch_device(), device=comfy.model_management.get_torch_device(),
dtype=comfy.model_management.VAE_DTYPE, dtype=comfy.model_management.unet_dtype(),
) )
return (controlnet,) return (controlnet,)
@@ -562,8 +693,8 @@ class DiffusersControlNetUnitStack:
def run( def run(
self, self,
controlnet_unit_1: tuple[ControlNetModel], controlnet_unit_1: tuple[ControlNetModel],
controlnet_unit_2: tuple[ControlNetModel] | None, controlnet_unit_2: tuple[ControlNetModel] | None = None,
controlnet_unit_3: tuple[ControlNetModel] | None, controlnet_unit_3: tuple[ControlNetModel] | None = None,
): ):
stack = [] stack = []
if controlnet_unit_1: if controlnet_unit_1:
@@ -623,13 +754,22 @@ class DiffusersGenerator:
"step": 64, "step": 64,
}, },
), ),
"reference_strength": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.01,
},
),
"reference_style_fidelity": ( "reference_style_fidelity": (
"FLOAT", "FLOAT",
{ {
"default": 0.5, "default": 0.5,
"min": 0.0, "min": 0.0,
"max": 1.0, "max": 1.0,
"step": 0.1, "step": 0.01,
}, },
), ),
}, },
@@ -675,6 +815,7 @@ class DiffusersGenerator:
reference_only_adain: str = "disable", reference_only_adain: str = "disable",
reference_image: torch.Tensor | None = None, reference_image: torch.Tensor | None = None,
reference_style_fidelity: float = 0.5, reference_style_fidelity: float = 0.5,
reference_strength: float = 1.0,
): ):
reference_only = reference_only == "enable" reference_only = reference_only == "enable"
reference_only_adain = reference_only_adain == "enable" reference_only_adain = reference_only_adain == "enable"
@@ -694,7 +835,7 @@ class DiffusersGenerator:
width=width, width=width,
generator=generator, generator=generator,
device=device, device=device,
dtype=comfy.model_management.VAE_DTYPE, dtype=comfy.model_management.vae_dtype(),
) )
images = latents_to_img_tensor(pipeline, latents) images = latents_to_img_tensor(pipeline, latents)
else: else:
@@ -730,6 +871,7 @@ class DiffusersGenerator:
strength=strength, strength=strength,
controlnet_units=controlnet_units, controlnet_units=controlnet_units,
callback=callback, callback=callback,
reference_strength=reference_strength,
reference_attn=reference_only, reference_attn=reference_only,
reference_adain=reference_only_adain, reference_adain=reference_only_adain,
style_fidelity=reference_style_fidelity, style_fidelity=reference_style_fidelity,
@@ -743,6 +885,8 @@ class DiffusersGenerator:
# 0 ~ 255 to 0 ~ 1 # 0 ~ 255 to 0 ~ 1
imgs = imgs / 255 imgs = imgs / 255
# (B, C, H, W) to (B, H, W, C) # (B, C, H, W) to (B, H, W, C)
torch.cuda.empty_cache()
gc.collect()
return (imgs,) return (imgs,)
@@ -750,6 +894,7 @@ NODE_CLASS_MAPPINGS = {
"GetFilledColorImage": GetFilledColorImage, "GetFilledColorImage": GetFilledColorImage,
"GetAverageColorFromImage": GetAverageColorFromImage, "GetAverageColorFromImage": GetAverageColorFromImage,
"DiffusersPipeline": DiffusersPipeline, "DiffusersPipeline": DiffusersPipeline,
"DiffusersXLPipeline": DiffusersXLPipeline,
"DiffusersGenerator": DiffusersGenerator, "DiffusersGenerator": DiffusersGenerator,
"DiffusersPrepareLatents": DiffusersPrepareLatents, "DiffusersPrepareLatents": DiffusersPrepareLatents,
"DiffusersDecoder": DiffusersDecoder, "DiffusersDecoder": DiffusersDecoder,
@@ -763,6 +908,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"GetFilledColorImage": "Get Filled Color Image Jannchie", "GetFilledColorImage": "Get Filled Color Image Jannchie",
"GetAverageColorFromImage": "Get Average Color From Image Jannchie", "GetAverageColorFromImage": "Get Average Color From Image Jannchie",
"DiffusersPipeline": "🤗 Diffusers Pipeline", "DiffusersPipeline": "🤗 Diffusers Pipeline",
"DiffusersXLPipeline": "🤗 Diffusers XL Pipeline",
"DiffusersGenerator": "🤗 Diffusers Generator", "DiffusersGenerator": "🤗 Diffusers Generator",
"DiffusersPrepareLatents": "🤗 Diffusers Prepare Latents", "DiffusersPrepareLatents": "🤗 Diffusers Prepare Latents",
"DiffusersDecoder": "🤗 Diffusers Decoder", "DiffusersDecoder": "🤗 Diffusers Decoder",
BIN
View File
Binary file not shown.

After

Width:  |  Height:  |  Size: 416 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 2.4 MiB

-475
View File
@@ -1,475 +0,0 @@
{
"last_node_id": 14,
"last_link_id": 33,
"nodes": [
{
"id": 1,
"type": "DiffusersPipeline",
"pos": [
110,
310
],
"size": {
"0": 370,
"1": 110
},
"flags": {},
"order": 0,
"mode": 0,
"outputs": [
{
"name": "pipeline",
"type": "DIFFUSERS_PIPELINE",
"links": [
2,
29
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "DiffusersPipeline"
},
"widgets_values": [
"picxReal_10.safetensors",
"vae-ft-mse-840000-ema-pruned.safetensors",
"-"
],
"shape": 1
},
{
"id": 3,
"type": "DiffusersCompelPromptEmbedding",
"pos": [
720,
80
],
"size": {
"0": 430.8000183105469,
"1": 200
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "pipeline",
"type": "DIFFUSERS_PIPELINE",
"link": 2
}
],
"outputs": [
{
"name": "positive prompt embedding",
"type": "DIFFUSERS_PROMPT_EMBEDDING",
"links": [
30
],
"shape": 3,
"slot_index": 0
},
{
"name": "negative prompt embedding",
"type": "DIFFUSERS_PROMPT_EMBEDDING",
"links": [
31
],
"shape": 3,
"slot_index": 1
}
],
"properties": {
"Node name for S&R": "DiffusersCompelPromptEmbedding"
},
"widgets_values": [
"(masterpiece)1.2, (best quality)1.4, ",
""
],
"shape": 1
},
{
"id": 10,
"type": "DiffusersControlnetLoader",
"pos": [
510,
480
],
"size": {
"0": 340,
"1": 60
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "controlnet",
"type": "DIFFUSERS_CONTROLNET",
"links": [
20
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "DiffusersControlnetLoader"
},
"widgets_values": [
"control_v11p_sd15_openpose_fp16.safetensors"
],
"shape": 1
},
{
"id": 9,
"type": "DiffusersControlnetUnit",
"pos": [
960,
480
],
"size": {
"0": 330,
"1": 126
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "controlnet",
"type": "DIFFUSERS_CONTROLNET",
"link": 20,
"slot_index": 0
},
{
"name": "image",
"type": "IMAGE",
"link": 18
}
],
"outputs": [
{
"name": "controlnet unit",
"type": "CONTROLNET_UNIT",
"links": [
32
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "DiffusersControlnetUnit"
},
"widgets_values": [
1,
0,
1
],
"shape": 1
},
{
"id": 8,
"type": "OpenposePreprocessor",
"pos": [
510,
620
],
"size": {
"0": 360,
"1": 150
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 17
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
18,
28
],
"shape": 3,
"slot_index": 0
},
{
"name": "POSE_KEYPOINT",
"type": "POSE_KEYPOINT",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "OpenposePreprocessor"
},
"widgets_values": [
"enable",
"enable",
"enable",
512
],
"shape": 1
},
{
"id": 13,
"type": "Image Comparer (rgthree)",
"pos": {
"0": 1880,
"1": 600,
"2": 0,
"3": 0,
"4": 0,
"5": 0,
"6": 0,
"7": 0,
"8": 0,
"9": 0
},
"size": {
"0": 603.9534301757812,
"1": 646.658935546875
},
"flags": {},
"order": 7,
"mode": 0,
"inputs": [
{
"name": "image_a",
"type": "IMAGE",
"link": 33,
"dir": 3
},
{
"name": "image_b",
"type": "IMAGE",
"link": 28,
"dir": 3
}
],
"outputs": [],
"properties": {
"comparer_mode": "Slide"
},
"widgets_values": [
[
"/view?filename=rgthree.compare._temp_pknnj_00003_.png&type=temp&subfolder=&rand=0.7251475711869597",
"/view?filename=rgthree.compare._temp_pknnj_00004_.png&type=temp&subfolder=&rand=0.5900207825118742"
]
],
"shape": 1
},
{
"id": 14,
"type": "DiffusersGenerator",
"pos": [
1380,
310
],
"size": {
"0": 405.5999755859375,
"1": 418
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "pipeline",
"type": "DIFFUSERS_PIPELINE",
"link": 29
},
{
"name": "positive_prompt_embedding",
"type": "DIFFUSERS_PROMPT_EMBEDDING",
"link": 30
},
{
"name": "negative_prompt_embedding",
"type": "DIFFUSERS_PROMPT_EMBEDDING",
"link": 31
},
{
"name": "images",
"type": "IMAGE",
"link": null
},
{
"name": "mask",
"type": "MASK",
"link": null
},
{
"name": "controlnet_units",
"type": "CONTROLNET_UNIT",
"link": 32
},
{
"name": "reference_image",
"type": "IMAGE",
"link": null
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
33
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "DiffusersGenerator"
},
"widgets_values": [
1,
30,
7,
127993911908,
"randomize",
1,
512,
512,
0.5,
"disable",
"disable"
]
},
{
"id": 7,
"type": "LoadImage",
"pos": [
110,
470
],
"size": {
"0": 370.4840393066406,
"1": 537.7867431640625
},
"flags": {},
"order": 2,
"mode": 0,
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
17
],
"shape": 3,
"slot_index": 0
},
{
"name": "MASK",
"type": "MASK",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"image_jjZSuko8_1701021290044_raw.webp",
"image"
],
"shape": 1
}
],
"links": [
[
2,
1,
0,
3,
0,
"DIFFUSERS_PIPELINE"
],
[
17,
7,
0,
8,
0,
"IMAGE"
],
[
18,
8,
0,
9,
1,
"IMAGE"
],
[
20,
10,
0,
9,
0,
"DIFFUSERS_CONTROLNET"
],
[
28,
8,
0,
13,
1,
"IMAGE"
],
[
29,
1,
0,
14,
0,
"DIFFUSERS_PIPELINE"
],
[
30,
3,
0,
14,
1,
"DIFFUSERS_PROMPT_EMBEDDING"
],
[
31,
3,
1,
14,
2,
"DIFFUSERS_PROMPT_EMBEDDING"
],
[
32,
9,
0,
14,
5,
"CONTROLNET_UNIT"
],
[
33,
14,
0,
13,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {},
"version": 0.4
}
Binary file not shown.

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 806 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 685 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

+35 -8
View File
@@ -1,4 +1,6 @@
from diffusers import AutoencoderKL, StableDiffusionImg2ImgPipeline import contextlib
from diffusers import AutoencoderKL, AutoencoderTiny, DPMSolverMultistepScheduler
from diffusers.schedulers import ( from diffusers.schedulers import (
DEISMultistepScheduler, DEISMultistepScheduler,
DPMSolverMultistepScheduler, DPMSolverMultistepScheduler,
@@ -13,6 +15,7 @@ from diffusers.schedulers import (
) )
import comfy.model_management import comfy.model_management
import folder_paths
from .jannchie import * from .jannchie import *
@@ -38,36 +41,60 @@ schedulers = {
"UniPC": UniPCMultistepScheduler(), "UniPC": UniPCMultistepScheduler(),
} }
class PipelineWrapper: class PipelineWrapper:
def __init__( def __init__(
self, ckpt_path: str, vae_path: str = None, scheduler_name: str = None self,
ckpt_path: str,
vae_path: str = None,
scheduler_name: str = None,
use_tiny_vae: bool = False,
): ):
scheduler = schedulers.get(scheduler_name) scheduler = schedulers.get(scheduler_name)
device = comfy.model_management.get_torch_device() device = comfy.model_management.get_torch_device()
dtype = comfy.model_management.VAE_DTYPE vae_dtype = comfy.model_management.vae_dtype()
unet_dtype = comfy.model_management.unet_dtype()
if ckpt_path.endswith(".safetensors"): if ckpt_path.endswith(".safetensors"):
self.pipeline = JannchiePipeline.from_single_file( self.pipeline = JannchiePipeline.from_single_file(
ckpt_path, ckpt_path,
torch_dtype=dtype, torch_dtype=unet_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers")[0],
use_safetensors=True,
) )
else: else:
self.pipeline = JannchiePipeline.from_pretrained( self.pipeline = JannchiePipeline.from_pretrained(
ckpt_path, ckpt_path,
torch_dtype=dtype, torch_dtype=unet_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers")[0],
use_safetensors=ckpt_path.endswith(".safetensors"),
) )
if vae_path:
if use_tiny_vae:
self.pipeline.vae = AutoencoderTiny.from_pretrained("madebyollin/taesd").to(
device=self.pipeline.device, dtype=vae_dtype
)
elif vae_path:
if vae_path.endswith(".safetensors"): if vae_path.endswith(".safetensors"):
self.pipeline.vae = AutoencoderKL.from_single_file( self.pipeline.vae = AutoencoderKL.from_single_file(
vae_path, vae_path,
torch_dtype=dtype, torch_dtype=vae_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers"),
use_safetensors=True,
) )
else: else:
self.pipeline.vae = AutoencoderKL.from_pretrained( self.pipeline.vae = AutoencoderKL.from_pretrained(
vae_path, vae_path,
torch_dtype=dtype, torch_dtype=vae_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers"),
use_safetensors=vae_path.endswith(".safetensors"),
) )
if scheduler: if scheduler:
self.pipeline.scheduler = scheduler self.pipeline.scheduler = scheduler
self.pipeline.to(device) self.pipeline.to(device)
self.pipeline.vae.to(vae_dtype)
self.pipeline.safety_checker = None self.pipeline.safety_checker = None
with contextlib.suppress(Exception):
self.pipeline.enable_xformers_memory_efficient_attention()
+104 -70
View File
@@ -35,7 +35,7 @@ from transformers import (
formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s") formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s")
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
logger.setLevel(logging.DEBUG) logger.setLevel(logging.INFO)
ch = logging.StreamHandler() ch = logging.StreamHandler()
ch.setFormatter(formatter) ch.setFormatter(formatter)
logger.addHandler(ch) logger.addHandler(ch)
@@ -225,10 +225,13 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
timesteps: List[int] = None, timesteps: List[int] = None,
mask_image: PipelineImageInput = None, mask_image: PipelineImageInput = None,
masked_image_latents: Optional[torch.FloatTensor] = None, masked_image_latents: Optional[torch.FloatTensor] = None,
ip_adapter_image: Optional[PipelineImageInput] = None,
ip_adapter_image_embeds: Optional[List[torch.FloatTensor]] = None,
reference_strength: float = 1.0,
*arg, *arg,
**args, **args,
): ):
device = self._execution_device device = self.unet.device
if height == None: if height == None:
if isinstance(image, torch.Tensor): if isinstance(image, torch.Tensor):
if image is not None: if image is not None:
@@ -400,7 +403,6 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
latent_timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt) latent_timestep = timesteps[:1].repeat(batch_size * num_images_per_prompt)
# 7. Prepare latent variables # 7. Prepare latent variables
logger.debug("Preparing latent variables")
num_channels_latents = self.unet.config.in_channels num_channels_latents = self.unet.config.in_channels
if image is not None: if image is not None:
if isinstance(image, PIL.Image.Image): if isinstance(image, PIL.Image.Image):
@@ -418,7 +420,6 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
False, # it will duplicate the latents after this step False, # it will duplicate the latents after this step
) )
logger.debug("Preparing latents")
num_channels_unet = self.unet.config.in_channels num_channels_unet = self.unet.config.in_channels
return_image_latents = num_channels_unet == 4 return_image_latents = num_channels_unet == 4
latents_outputs = self.prepare_latents( latents_outputs = self.prepare_latents(
@@ -445,9 +446,9 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
if mask_image is not None: if mask_image is not None:
mask_condition = self.mask_processor.preprocess( mask_condition = self.mask_processor.preprocess(
mask_image, height=height, width=width mask_image, height=height, width=width
) ).to(device=device)
init_image = image init_image = image
init_image = init_image.to(dtype=torch.float32) init_image = init_image.to(dtype=torch.float32, device=device)
if masked_image_latents is None: if masked_image_latents is None:
masked_image = init_image * (mask_condition < 0.5) masked_image = init_image * (mask_condition < 0.5)
else: else:
@@ -458,7 +459,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
batch_size * num_images_per_prompt, batch_size * num_images_per_prompt,
height, height,
width, width,
prompt_embeds.dtype, self.unet.dtype,
device, device,
generator, generator,
do_classifier_free_guidance, do_classifier_free_guidance,
@@ -478,6 +479,25 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
# 9. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline # 9. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta) extra_step_kwargs = self.prepare_extra_step_kwargs(generator, eta)
if ip_adapter_image is not None or ip_adapter_image_embeds is not None:
image_embeds = self.prepare_ip_adapter_image_embeds(
ip_adapter_image,
ip_adapter_image_embeds,
device,
batch_size * num_images_per_prompt,
do_classifier_free_guidance,
)
# Add image embeds for IP-Adapter
added_cond_kwargs = (
{"image_embeds": image_embeds}
if (ip_adapter_image is not None or ip_adapter_image_embeds is not None)
else {}
)
# text_embeds for reference, TODO: I forgot why it is needed
added_cond_kwargs["text_embeds"] = prompt_embeds
ref_mask_dict, out_mask_dict = self.get_ref_mask_dicts( ref_mask_dict, out_mask_dict = self.get_ref_mask_dicts(
ref_image_mask, ref_image_mask,
height, height,
@@ -496,6 +516,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
gn_auto_machine_weight=gn_auto_machine_weight, gn_auto_machine_weight=gn_auto_machine_weight,
ref_mask_dict=ref_mask_dict, ref_mask_dict=ref_mask_dict,
out_mask_dict=out_mask_dict, out_mask_dict=out_mask_dict,
strength=reference_strength,
) )
if reference_attn: if reference_attn:
self.unet = ReferenceOnlyUNet2DConditionModel.from_unet( self.unet = ReferenceOnlyUNet2DConditionModel.from_unet(
@@ -566,6 +587,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
encoder_hidden_states=prompt_embeds, encoder_hidden_states=prompt_embeds,
cross_attention_kwargs=cross_attention_kwargs, cross_attention_kwargs=cross_attention_kwargs,
return_dict=False, return_dict=False,
added_cond_kwargs=added_cond_kwargs,
) )
self.unet.ref_data.MODE = "read" self.unet.ref_data.MODE = "read"
@@ -607,7 +629,9 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
) )
if n_controlnet_unit != 0: if n_controlnet_unit != 0:
down_block_res_samples, mid_block_res_sample = self.controlnet( down_block_res_samples, mid_block_res_sample = self.controlnet(
control_model_input, control_model_input.to(
device=device, dtype=self.controlnet.dtype
),
t, t,
encoder_hidden_states=controlnet_prompt_embeds, encoder_hidden_states=controlnet_prompt_embeds,
controlnet_cond=controlnet_images, controlnet_cond=controlnet_images,
@@ -634,12 +658,15 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
down_block_res_samples, mid_block_res_sample = None, None down_block_res_samples, mid_block_res_sample = None, None
# predict the noise residual # predict the noise residual
noise_pred = self.unet( noise_pred = self.unet(
latent_model_input, latent_model_input.to(device=device, dtype=self.unet.dtype),
t, t,
encoder_hidden_states=prompt_embeds, encoder_hidden_states=prompt_embeds.to(
device=device, dtype=self.unet.dtype
),
cross_attention_kwargs=cross_attention_kwargs, cross_attention_kwargs=cross_attention_kwargs,
down_block_additional_residuals=down_block_res_samples, down_block_additional_residuals=down_block_res_samples,
mid_block_additional_residual=mid_block_res_sample, mid_block_additional_residual=mid_block_res_sample,
added_cond_kwargs=added_cond_kwargs,
)["sample"] )["sample"]
# perform guidance # perform guidance
if do_classifier_free_guidance: if do_classifier_free_guidance:
@@ -666,6 +693,9 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
init_latents_proper = self.scheduler.add_noise( init_latents_proper = self.scheduler.add_noise(
init_latents_proper, noise, torch.tensor([noise_timestep]) init_latents_proper, noise, torch.tensor([noise_timestep])
) )
init_latents_proper = init_latents_proper.to(
device=device, dtype=self.unet.dtype
)
input_latents = ( input_latents = (
1 - init_mask 1 - init_mask
@@ -686,6 +716,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
torch.cuda.empty_cache() torch.cuda.empty_cache()
if output_type != "latent": if output_type != "latent":
input_latents = input_latents.to(device=device, dtype=self.vae.dtype)
result_imgs = self.vae.decode( result_imgs = self.vae.decode(
input_latents / self.vae.config.scaling_factor, return_dict=False input_latents / self.vae.config.scaling_factor, return_dict=False
)[0] )[0]
@@ -842,6 +873,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
self.mask_processor = VaeImageProcessor( self.mask_processor = VaeImageProcessor(
vae_scale_factor=self.vae_scale_factor, vae_scale_factor=self.vae_scale_factor,
do_normalize=False, do_normalize=False,
do_resize=False,
do_binarize=True, do_binarize=True,
do_convert_grayscale=True, do_convert_grayscale=True,
) )
@@ -860,16 +892,13 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
# encode the mask image into latents space so we can concatenate it to the latents # encode the mask image into latents space so we can concatenate it to the latents
if isinstance(generator, list): if isinstance(generator, list):
image_latents = [ image_latents = [
self.vae.encode(image[i : i + 1]).latent_dist.sample( retrieve_latents(self.vae.encode(image[i : i + 1]), generator[i])
generator=generator[i]
)
for i in range(batch_size) for i in range(batch_size)
] ]
image_latents = torch.cat(image_latents, dim=0) image_latents = torch.cat(image_latents, dim=0)
else: else:
image_latents = self.vae.encode(image).latent_dist.sample( image = image.to(self.vae.dtype)
generator=generator image_latents = retrieve_latents(self.vae.encode(image), generator)
)
image_latents = self.vae.config.scaling_factor * image_latents image_latents = self.vae.config.scaling_factor * image_latents
# duplicate mask and ref_image_latents for each generation per prompt, using mps friendly method # duplicate mask and ref_image_latents for each generation per prompt, using mps friendly method
@@ -930,17 +959,15 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
if isinstance(generator, list): if isinstance(generator, list):
image_latent = torch.cat( image_latent = torch.cat(
[ [
self.vae.encode(image_tensor[i : i + 1]).latent_dist.sample( retrieve_latents(
generator=generator[i] self.vae.encode(image_tensor[i : i + 1]), generator[i]
) )
for i in range(image_tensor.shape[0]) for i in range(image_tensor.shape[0])
], ],
dim=0, dim=0,
) )
else: else:
image_latent = self.vae.encode(image_tensor).latent_dist.sample( image_latent = retrieve_latents(self.vae.encode(image_tensor), generator)
generator=generator
)
image_latent = self.vae.config.scaling_factor * image_latent image_latent = self.vae.config.scaling_factor * image_latent
return image_latent.to(device=device, dtype=dtype) return image_latent.to(device=device, dtype=dtype)
@@ -979,6 +1006,8 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
) )
if return_image_latents or (latents is None and not is_strength_max): if return_image_latents or (latents is None and not is_strength_max):
# TODO: check it # TODO: check it
if image is None:
image = torch.randn(shape, device=device, dtype=dtype)
image = image.to(device=device, dtype=dtype) image = image.to(device=device, dtype=dtype)
if image.shape[1] == 4: if image.shape[1] == 4:
@@ -988,6 +1017,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
image_latents = image_latents.repeat( image_latents = image_latents.repeat(
batch_size // image_latents.shape[0], 1, 1, 1 batch_size // image_latents.shape[0], 1, 1, 1
) )
image_latents.to(device=device, dtype=dtype)
if latents is None: if latents is None:
noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype) noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
@@ -1081,6 +1111,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
] ]
image_latents = torch.cat(image_latents, dim=0) image_latents = torch.cat(image_latents, dim=0)
else: else:
image = image.to(self.vae.dtype)
image_latents = retrieve_latents( image_latents = retrieve_latents(
self.vae.encode(image), generator=generator self.vae.encode(image), generator=generator
) )
@@ -1103,6 +1134,7 @@ class ReferenceData:
gn_auto_machine_weight: float = 1.0 gn_auto_machine_weight: float = 1.0
ref_mask_dict: dict = None ref_mask_dict: dict = None
out_mask_dict: dict = None out_mask_dict: dict = None
strength: float = 1.0
class ReferenceOnlyUNet2DConditionModel(UNet2DConditionModel): class ReferenceOnlyUNet2DConditionModel(UNet2DConditionModel):
@@ -1211,10 +1243,6 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
bank = self.bank bank = self.bank
assert isinstance(bank, list) assert isinstance(bank, list)
uc_mask = ref_data.uc_mask
bool_mask = ref_data.bool_mask
ref_mask_dict = ref_data.ref_mask_dict
if self.use_ada_layer_norm: if self.use_ada_layer_norm:
norm_hidden_states = self.norm1(hidden_states, timestep) norm_hidden_states = self.norm1(hidden_states, timestep)
elif self.use_ada_layer_norm_zero: elif self.use_ada_layer_norm_zero:
@@ -1260,9 +1288,35 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
attention_mask=attention_mask, attention_mask=attention_mask,
**cross_attention_kwargs, **cross_attention_kwargs,
) )
else: elif ref_data.MODE == "read":
if ref_data.MODE == "write": style_fidelity = ref_data.style_fidelity
bank.append(norm_hidden_states.detach().clone()) attention_auto_machine_weight = ref_data.attention_auto_machine_weight
if attention_auto_machine_weight > self.attn_weight:
attn_output_uc = self.attn1(
norm_hidden_states,
encoder_hidden_states=torch.cat(
[norm_hidden_states] + self.bank, dim=1
),
# attention_mask=attention_mask,
**cross_attention_kwargs,
)
attn_output_c = attn_output_uc.clone()
do_classifier_free_guidance = ref_data.do_classifier_free_guidance
if do_classifier_free_guidance and style_fidelity > 0:
uc_mask = ref_data.uc_mask
attn_output_c[uc_mask] = self.attn1(
norm_hidden_states[uc_mask],
encoder_hidden_states=norm_hidden_states[uc_mask],
**cross_attention_kwargs,
)
attn_output = (
style_fidelity * attn_output_c
+ (1.0 - style_fidelity) * attn_output_uc
)
attn_output *= ref_data.strength
bank.clear()
else:
# without reference only
attn_output = self.attn1( attn_output = self.attn1(
norm_hidden_states, norm_hidden_states,
encoder_hidden_states=( encoder_hidden_states=(
@@ -1271,42 +1325,17 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
attention_mask=attention_mask, attention_mask=attention_mask,
**cross_attention_kwargs, **cross_attention_kwargs,
) )
if ref_data.MODE == "read":
style_fidelity = ref_data.style_fidelity
attention_auto_machine_weight = ref_data.attention_auto_machine_weight
do_classifier_free_guidance = ref_data.do_classifier_free_guidance
if attention_auto_machine_weight > self.attn_weight:
attn_output_uc = self.attn1(
norm_hidden_states,
encoder_hidden_states=torch.cat(
[norm_hidden_states] + self.bank, dim=1
),
# attention_mask=attention_mask,
**cross_attention_kwargs,
)
attn_output_c = attn_output_uc.clone()
if do_classifier_free_guidance and style_fidelity > 0:
attn_output_c[uc_mask] = self.attn1(
norm_hidden_states[uc_mask],
encoder_hidden_states=norm_hidden_states[uc_mask],
**cross_attention_kwargs,
)
attn_output = (
style_fidelity * attn_output_c
+ (1.0 - style_fidelity) * attn_output_uc
)
bank.clear()
else:
# 原始的自注意力(无 reference only
attn_output = self.attn1(
norm_hidden_states,
encoder_hidden_states=(
encoder_hidden_states if self.only_cross_attention else None
),
attention_mask=attention_mask,
**cross_attention_kwargs,
)
elif ref_data.MODE == "write":
bank.append(norm_hidden_states.detach().clone())
attn_output = self.attn1(
norm_hidden_states,
encoder_hidden_states=(
encoder_hidden_states if self.only_cross_attention else None
),
attention_mask=attention_mask,
**cross_attention_kwargs,
)
if self.use_ada_layer_norm_zero: if self.use_ada_layer_norm_zero:
attn_output = gate_msa.unsqueeze(1) * attn_output attn_output = gate_msa.unsqueeze(1) * attn_output
@@ -1315,7 +1344,6 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
# 2.5 GLIGEN Control # 2.5 GLIGEN Control
if gligen_kwargs is not None: if gligen_kwargs is not None:
hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"]) hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"])
# 2.5 ends
# 2. Cross-Attention # 2. Cross-Attention
if self.attn2 is not None: if self.attn2 is not None:
@@ -1348,9 +1376,7 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
if self.use_ada_layer_norm_zero: if self.use_ada_layer_norm_zero:
ff_output = gate_mlp.unsqueeze(1) * ff_output ff_output = gate_mlp.unsqueeze(1) * ff_output
hidden_states = ff_output + hidden_states return ff_output + hidden_states
return hidden_states
class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D): class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
@@ -1381,7 +1407,8 @@ class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
# TODO(Patrick, William) - attention mask is not used # TODO(Patrick, William) - attention mask is not used
output_states = () output_states = ()
for i, (resnet, attn) in enumerate(zip(self.resnets, self.attentions)): blocks = list(zip(self.resnets, self.attentions))
for i, (resnet, attn) in enumerate(blocks):
hidden_states = resnet(hidden_states, temb) hidden_states = resnet(hidden_states, temb)
hidden_states = attn( hidden_states = attn(
hidden_states, hidden_states,
@@ -1413,7 +1440,10 @@ class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
style_fidelity * hidden_states_c style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc + (1.0 - style_fidelity) * hidden_states_uc
) )
hidden_states *= self.ref_data.strength
# apply additional residuals to the output of the last pair of resnet and attention blocks
if i == len(blocks) - 1 and additional_residuals is not None:
hidden_states = hidden_states + additional_residuals
output_states = output_states + (hidden_states,) output_states = output_states + (hidden_states,)
if MODE == "read": if MODE == "read":
@@ -1476,6 +1506,7 @@ class DownBlock2DReferenceOnly(DownBlock2D):
style_fidelity * hidden_states_c style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc + (1.0 - style_fidelity) * hidden_states_uc
) )
hidden_states *= self.ref_data.strength
output_states = output_states + (hidden_states,) output_states = output_states + (hidden_states,)
@@ -1523,7 +1554,6 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
do_classifier_free_guidance = self.ref_data.do_classifier_free_guidance do_classifier_free_guidance = self.ref_data.do_classifier_free_guidance
style_fidelity = self.ref_data.style_fidelity style_fidelity = self.ref_data.style_fidelity
uc_mask = self.ref_data.uc_mask uc_mask = self.ref_data.uc_mask
eps = 1e-6
x = super().forward(*args, **kwargs) x = super().forward(*args, **kwargs)
if MODE == "write" and gn_auto_machine_weight >= self.gn_weight: if MODE == "write" and gn_auto_machine_weight >= self.gn_weight:
var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
@@ -1532,6 +1562,7 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
if MODE == "read": if MODE == "read":
if len(self.mean_bank) > 0 and len(self.var_bank) > 0: if len(self.mean_bank) > 0 and len(self.var_bank) > 0:
var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0) var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
eps = 1e-6
std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5 std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5
mean_acc = sum(self.mean_bank) / float(len(self.mean_bank)) mean_acc = sum(self.mean_bank) / float(len(self.mean_bank))
var_acc = sum(self.var_bank) / float(len(self.var_bank)) var_acc = sum(self.var_bank) / float(len(self.var_bank))
@@ -1541,6 +1572,7 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
if do_classifier_free_guidance and style_fidelity > 0: if do_classifier_free_guidance and style_fidelity > 0:
x_c[uc_mask] = x[uc_mask] x_c[uc_mask] = x[uc_mask]
x = style_fidelity * x_c + (1.0 - style_fidelity) * x_uc x = style_fidelity * x_c + (1.0 - style_fidelity) * x_uc
x *= self.ref_data.strength
self.mean_bank = [] self.mean_bank = []
self.var_bank = [] self.var_bank = []
return x return x
@@ -1598,6 +1630,7 @@ class UpBlock2DReferenceOnly(UpBlock2D):
style_fidelity * hidden_states_c style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc + (1.0 - style_fidelity) * hidden_states_uc
) )
hidden_states *= self.ref_data.strength
if MODE == "read": if MODE == "read":
self.mean_bank = [] self.mean_bank = []
@@ -1674,6 +1707,7 @@ class CrossAttnUpBlock2DReferenceOnly(CrossAttnUpBlock2D):
style_fidelity * hidden_states_c style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc + (1.0 - style_fidelity) * hidden_states_uc
) )
hidden_states *= self.ref_data.strength
if MODE == "read": if MODE == "read":
self.mean_bank = [] self.mean_bank = []
+15
View File
@@ -0,0 +1,15 @@
[project]
name = "comfyui-j"
description = "This is a completely different set of nodes than Comfy's own KSampler series. This set of nodes is based on Diffusers, which makes it easier to import models, apply prompts with weights, inpaint, reference only, controlnet, etc."
version = "1.1.0"
license = "LICENSE"
dependencies = ["compel", "diffusers", "numpy"]
[project.urls]
Repository = "https://github.com/Jannchie/ComfyUI-J"
# Used by Comfy Registry https://comfyregistry.org
[tool.comfy]
PublisherId = ""
DisplayName = "ComfyUI-J"
Icon = ""