20 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
12 changed files with 375 additions and 125 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 }}
+6
View File
@@ -33,6 +33,8 @@ ref_only supports two modes: attn and attn + adain, and can adjust the style fid
### 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
@@ -58,6 +60,10 @@ A checkpoint for stablediffusion 1.5 is all your need. But for full automation,
![Change Clothes](./examples/change_clothes.png)
## QR Code
![QR Code](./examples/qr_code.png)
## FAQ
### Why Diffusers?
+195 -49
View File
@@ -1,4 +1,5 @@
import contextlib
import gc
import random
from collections import Counter
@@ -74,6 +75,9 @@ def latents_to_img_tensor(pipeline, latents):
# 1. 输入的 latents 是一个 -1 ~ 1 之间的 tensor
# 2. 先进行缩放
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
dec_tensor = pipeline.vae.decode(scaled_latents, return_dict=False)[0]
# 4. 缩放到 0 ~ 1 之间
@@ -197,7 +201,7 @@ class GetFilledColorImage:
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
"display": "number",
},
),
@@ -207,7 +211,7 @@ class GetFilledColorImage:
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
"display": "number",
},
),
@@ -217,7 +221,7 @@ class GetFilledColorImage:
"default": 0.0,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
"display": "number",
},
),
@@ -283,6 +287,7 @@ class DiffusersTextureInversionLoader:
path = folder_paths.get_full_path("embeddings", texture_inversion)
token = texture_inversion.split(".")[0]
pipeline.load_textual_inversion(path, token=token)
print(f"Loaded {texture_inversion}")
return (pipeline,)
@@ -297,7 +302,7 @@ class GetAverageColorFromImage:
return {
"required": {
"image": ("IMAGE",),
"average": ("STRING", {"default": "mean", "options": ["mean", "mode"]}),
"average": (("mean", "mode"),),
},
"optional": {
"mask": ("MASK",),
@@ -305,48 +310,112 @@ class GetAverageColorFromImage:
}
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":
return self.run_avg(image, mask)
elif average == "mode":
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):
if mask is not None:
mask = mask.unsqueeze(1)
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)
return tuple(average_rgb.squeeze().tolist())
pixel_sum = torch.sum(masked_image, dim=(1, 2))
if mask is not None:
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):
image = image.permute(0, 3, 1, 2)
if mask is not None:
mask = mask.unsqueeze(1)
image = image * mask
masked_image = image * mask if mask is not None else image
pixel_values = masked_image.view(
masked_image.shape[0], masked_image.shape[1], -1
# Flatten the image to a 2D matrix where each row is a color
flattened_image = image.view(-1, image.shape[-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)
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)
return (self.pipeline_wrapper.pipeline,)
class DiffusersPipeline:
@@ -372,11 +441,27 @@ class DiffusersPipeline:
"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)
if ckpt_path is None:
ckpt_path = ckpt_name
if vae_name == "-":
vae_path = None
else:
@@ -384,7 +469,9 @@ class DiffusersPipeline:
if scheduler_name == "-":
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,)
@@ -431,7 +518,7 @@ class DiffusersPrepareLatents:
batch_size=batch_size,
height=height,
width=width,
dtype=comfy.model_management.VAE_DTYPE,
dtype=comfy.model_management.vae_dtype(),
device=device,
generator=generator,
latents=latents,
@@ -459,6 +546,27 @@ class DiffusersDecoder:
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:
CATEGORY = "Jannchie"
FUNCTION = "run"
@@ -469,19 +577,42 @@ class DiffusersControlNetLoader:
def INPUT_TYPES(cls):
return {
"required": {
"controlnet_model_name": (
folder_paths.get_filename_list("controlnet"),
),
"controlnet_model_name": (controlnet_list,),
},
"optional": {
"controlnet_model_file": (folder_paths.get_filename_list("controlnet"),)
},
}
def run(self, controlnet_model_name: str):
controlnet_model_path = folder_paths.get_full_path(
"controlnet", controlnet_model_name
)
controlnet = ControlNetModel.from_single_file(controlnet_model_path).to(
def run(self, controlnet_model_name: str, controlnet_model_file: str = ""):
file_list = folder_paths.get_filename_list("controlnet")
if controlnet_model_name == "other":
controlnet_model_path = folder_paths.get_full_path(
"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(),
dtype=comfy.model_management.VAE_DTYPE,
dtype=comfy.model_management.unet_dtype(),
)
return (controlnet,)
@@ -562,8 +693,8 @@ class DiffusersControlNetUnitStack:
def run(
self,
controlnet_unit_1: tuple[ControlNetModel],
controlnet_unit_2: tuple[ControlNetModel] | None,
controlnet_unit_3: tuple[ControlNetModel] | None,
controlnet_unit_2: tuple[ControlNetModel] | None = None,
controlnet_unit_3: tuple[ControlNetModel] | None = None,
):
stack = []
if controlnet_unit_1:
@@ -623,13 +754,22 @@ class DiffusersGenerator:
"step": 64,
},
),
"reference_strength": (
"FLOAT",
{
"default": 1.0,
"min": 0.0,
"max": 1.0,
"step": 0.01,
},
),
"reference_style_fidelity": (
"FLOAT",
{
"default": 0.5,
"min": 0.0,
"max": 1.0,
"step": 0.1,
"step": 0.01,
},
),
},
@@ -675,6 +815,7 @@ class DiffusersGenerator:
reference_only_adain: str = "disable",
reference_image: torch.Tensor | None = None,
reference_style_fidelity: float = 0.5,
reference_strength: float = 1.0,
):
reference_only = reference_only == "enable"
reference_only_adain = reference_only_adain == "enable"
@@ -694,7 +835,7 @@ class DiffusersGenerator:
width=width,
generator=generator,
device=device,
dtype=comfy.model_management.VAE_DTYPE,
dtype=comfy.model_management.vae_dtype(),
)
images = latents_to_img_tensor(pipeline, latents)
else:
@@ -730,6 +871,7 @@ class DiffusersGenerator:
strength=strength,
controlnet_units=controlnet_units,
callback=callback,
reference_strength=reference_strength,
reference_attn=reference_only,
reference_adain=reference_only_adain,
style_fidelity=reference_style_fidelity,
@@ -743,6 +885,8 @@ class DiffusersGenerator:
# 0 ~ 255 to 0 ~ 1
imgs = imgs / 255
# (B, C, H, W) to (B, H, W, C)
torch.cuda.empty_cache()
gc.collect()
return (imgs,)
@@ -750,6 +894,7 @@ NODE_CLASS_MAPPINGS = {
"GetFilledColorImage": GetFilledColorImage,
"GetAverageColorFromImage": GetAverageColorFromImage,
"DiffusersPipeline": DiffusersPipeline,
"DiffusersXLPipeline": DiffusersXLPipeline,
"DiffusersGenerator": DiffusersGenerator,
"DiffusersPrepareLatents": DiffusersPrepareLatents,
"DiffusersDecoder": DiffusersDecoder,
@@ -763,6 +908,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"GetFilledColorImage": "Get Filled Color Image Jannchie",
"GetAverageColorFromImage": "Get Average Color From Image Jannchie",
"DiffusersPipeline": "🤗 Diffusers Pipeline",
"DiffusersXLPipeline": "🤗 Diffusers XL Pipeline",
"DiffusersGenerator": "🤗 Diffusers Generator",
"DiffusersPrepareLatents": "🤗 Diffusers Prepare Latents",
"DiffusersDecoder": "🤗 Diffusers Decoder",
BIN
View File
Binary file not shown.

Before

Width:  |  Height:  |  Size: 362 KiB

After

Width:  |  Height:  |  Size: 416 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 2.5 MiB

After

Width:  |  Height:  |  Size: 2.4 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 1.1 MiB

After

Width:  |  Height:  |  Size: 1.2 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 792 KiB

After

Width:  |  Height:  |  Size: 806 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 1.7 MiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 687 KiB

After

Width:  |  Height:  |  Size: 685 KiB

+32 -9
View File
@@ -1,7 +1,6 @@
import contextlib
import comfy.model_management
from diffusers import AutoencoderKL
from diffusers import AutoencoderKL, AutoencoderTiny, DPMSolverMultistepScheduler
from diffusers.schedulers import (
DEISMultistepScheduler,
DPMSolverMultistepScheduler,
@@ -15,6 +14,9 @@ from diffusers.schedulers import (
UniPCMultistepScheduler,
)
import comfy.model_management
import folder_paths
from .jannchie import *
schedulers = {
@@ -43,35 +45,56 @@ schedulers = {
class PipelineWrapper:
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)
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"):
self.pipeline = JannchiePipeline.from_single_file(
ckpt_path,
torch_dtype=dtype,
torch_dtype=unet_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers")[0],
use_safetensors=True,
)
else:
self.pipeline = JannchiePipeline.from_pretrained(
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"):
self.pipeline.vae = AutoencoderKL.from_single_file(
vae_path,
torch_dtype=dtype,
torch_dtype=vae_dtype,
cache_dir=folder_paths.get_folder_paths("diffusers"),
use_safetensors=True,
)
else:
self.pipeline.vae = AutoencoderKL.from_pretrained(
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:
self.pipeline.scheduler = scheduler
self.pipeline.to(device)
self.pipeline.vae.to(vae_dtype)
self.pipeline.safety_checker = None
with contextlib.suppress(Exception):
self.pipeline.enable_xformers_memory_efficient_attention()
+102 -67
View File
@@ -225,10 +225,13 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
timesteps: List[int] = None,
mask_image: PipelineImageInput = 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,
**args,
):
device = self._execution_device
device = self.unet.device
if height == None:
if isinstance(image, torch.Tensor):
if image is not None:
@@ -441,14 +444,12 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
# 7. Prepare mask latent variables
if mask_image is not None:
print(height, width)
mask_condition = self.mask_processor.preprocess(
mask_image, height=height, width=width
)
).to(device=device)
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:
print(init_image.shape, mask_condition.shape)
masked_image = init_image * (mask_condition < 0.5)
else:
masked_image = masked_image_latents
@@ -458,7 +459,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
batch_size * num_images_per_prompt,
height,
width,
prompt_embeds.dtype,
self.unet.dtype,
device,
generator,
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
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_image_mask,
height,
@@ -496,6 +516,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
gn_auto_machine_weight=gn_auto_machine_weight,
ref_mask_dict=ref_mask_dict,
out_mask_dict=out_mask_dict,
strength=reference_strength,
)
if reference_attn:
self.unet = ReferenceOnlyUNet2DConditionModel.from_unet(
@@ -566,6 +587,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
encoder_hidden_states=prompt_embeds,
cross_attention_kwargs=cross_attention_kwargs,
return_dict=False,
added_cond_kwargs=added_cond_kwargs,
)
self.unet.ref_data.MODE = "read"
@@ -607,7 +629,9 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
)
if n_controlnet_unit != 0:
down_block_res_samples, mid_block_res_sample = self.controlnet(
control_model_input,
control_model_input.to(
device=device, dtype=self.controlnet.dtype
),
t,
encoder_hidden_states=controlnet_prompt_embeds,
controlnet_cond=controlnet_images,
@@ -634,12 +658,15 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
down_block_res_samples, mid_block_res_sample = None, None
# predict the noise residual
noise_pred = self.unet(
latent_model_input,
latent_model_input.to(device=device, dtype=self.unet.dtype),
t,
encoder_hidden_states=prompt_embeds,
encoder_hidden_states=prompt_embeds.to(
device=device, dtype=self.unet.dtype
),
cross_attention_kwargs=cross_attention_kwargs,
down_block_additional_residuals=down_block_res_samples,
mid_block_additional_residual=mid_block_res_sample,
added_cond_kwargs=added_cond_kwargs,
)["sample"]
# perform guidance
if do_classifier_free_guidance:
@@ -666,6 +693,9 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
init_latents_proper = self.scheduler.add_noise(
init_latents_proper, noise, torch.tensor([noise_timestep])
)
init_latents_proper = init_latents_proper.to(
device=device, dtype=self.unet.dtype
)
input_latents = (
1 - init_mask
@@ -686,6 +716,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
torch.cuda.empty_cache()
if output_type != "latent":
input_latents = input_latents.to(device=device, dtype=self.vae.dtype)
result_imgs = self.vae.decode(
input_latents / self.vae.config.scaling_factor, return_dict=False
)[0]
@@ -861,16 +892,13 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
# encode the mask image into latents space so we can concatenate it to the latents
if isinstance(generator, list):
image_latents = [
self.vae.encode(image[i : i + 1]).latent_dist.sample(
generator=generator[i]
)
retrieve_latents(self.vae.encode(image[i : i + 1]), generator[i])
for i in range(batch_size)
]
image_latents = torch.cat(image_latents, dim=0)
else:
image_latents = self.vae.encode(image).latent_dist.sample(
generator=generator
)
image = image.to(self.vae.dtype)
image_latents = retrieve_latents(self.vae.encode(image), generator)
image_latents = self.vae.config.scaling_factor * image_latents
# duplicate mask and ref_image_latents for each generation per prompt, using mps friendly method
@@ -931,17 +959,15 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
if isinstance(generator, list):
image_latent = torch.cat(
[
self.vae.encode(image_tensor[i : i + 1]).latent_dist.sample(
generator=generator[i]
retrieve_latents(
self.vae.encode(image_tensor[i : i + 1]), generator[i]
)
for i in range(image_tensor.shape[0])
],
dim=0,
)
else:
image_latent = self.vae.encode(image_tensor).latent_dist.sample(
generator=generator
)
image_latent = retrieve_latents(self.vae.encode(image_tensor), generator)
image_latent = self.vae.config.scaling_factor * image_latent
return image_latent.to(device=device, dtype=dtype)
@@ -980,6 +1006,8 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
)
if return_image_latents or (latents is None and not is_strength_max):
# TODO: check it
if image is None:
image = torch.randn(shape, device=device, dtype=dtype)
image = image.to(device=device, dtype=dtype)
if image.shape[1] == 4:
@@ -989,6 +1017,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
image_latents = image_latents.repeat(
batch_size // image_latents.shape[0], 1, 1, 1
)
image_latents.to(device=device, dtype=dtype)
if latents is None:
noise = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
@@ -1082,6 +1111,7 @@ class JannchiePipeline(StableDiffusionControlNetPipeline):
]
image_latents = torch.cat(image_latents, dim=0)
else:
image = image.to(self.vae.dtype)
image_latents = retrieve_latents(
self.vae.encode(image), generator=generator
)
@@ -1104,6 +1134,7 @@ class ReferenceData:
gn_auto_machine_weight: float = 1.0
ref_mask_dict: dict = None
out_mask_dict: dict = None
strength: float = 1.0
class ReferenceOnlyUNet2DConditionModel(UNet2DConditionModel):
@@ -1212,8 +1243,6 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
bank = self.bank
assert isinstance(bank, list)
uc_mask = ref_data.uc_mask
if self.use_ada_layer_norm:
norm_hidden_states = self.norm1(hidden_states, timestep)
elif self.use_ada_layer_norm_zero:
@@ -1259,9 +1288,35 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
attention_mask=attention_mask,
**cross_attention_kwargs,
)
else:
if ref_data.MODE == "write":
bank.append(norm_hidden_states.detach().clone())
elif ref_data.MODE == "read":
style_fidelity = ref_data.style_fidelity
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(
norm_hidden_states,
encoder_hidden_states=(
@@ -1270,42 +1325,17 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
attention_mask=attention_mask,
**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:
attn_output = gate_msa.unsqueeze(1) * attn_output
@@ -1314,7 +1344,6 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
# 2.5 GLIGEN Control
if gligen_kwargs is not None:
hidden_states = self.fuser(hidden_states, gligen_kwargs["objs"])
# 2.5 ends
# 2. Cross-Attention
if self.attn2 is not None:
@@ -1347,9 +1376,7 @@ class BasicTransformerBlockReferenceOnly(BasicTransformerBlock):
if self.use_ada_layer_norm_zero:
ff_output = gate_mlp.unsqueeze(1) * ff_output
hidden_states = ff_output + hidden_states
return hidden_states
return ff_output + hidden_states
class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
@@ -1380,7 +1407,8 @@ class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
# TODO(Patrick, William) - attention mask is not used
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 = attn(
hidden_states,
@@ -1412,7 +1440,10 @@ class CrossAttnDownBlock2DReferenceOnly(CrossAttnDownBlock2D):
style_fidelity * hidden_states_c
+ (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,)
if MODE == "read":
@@ -1475,6 +1506,7 @@ class DownBlock2DReferenceOnly(DownBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
output_states = output_states + (hidden_states,)
@@ -1522,7 +1554,6 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
do_classifier_free_guidance = self.ref_data.do_classifier_free_guidance
style_fidelity = self.ref_data.style_fidelity
uc_mask = self.ref_data.uc_mask
eps = 1e-6
x = super().forward(*args, **kwargs)
if MODE == "write" and gn_auto_machine_weight >= self.gn_weight:
var, mean = torch.var_mean(x, dim=(2, 3), keepdim=True, correction=0)
@@ -1531,6 +1562,7 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
if MODE == "read":
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)
eps = 1e-6
std = torch.maximum(var, torch.zeros_like(var) + eps) ** 0.5
mean_acc = sum(self.mean_bank) / float(len(self.mean_bank))
var_acc = sum(self.var_bank) / float(len(self.var_bank))
@@ -1540,6 +1572,7 @@ class UNetMidBlock2DCrossAttnReferenceOnly(UNetMidBlock2DCrossAttn):
if do_classifier_free_guidance and style_fidelity > 0:
x_c[uc_mask] = x[uc_mask]
x = style_fidelity * x_c + (1.0 - style_fidelity) * x_uc
x *= self.ref_data.strength
self.mean_bank = []
self.var_bank = []
return x
@@ -1597,6 +1630,7 @@ class UpBlock2DReferenceOnly(UpBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
if MODE == "read":
self.mean_bank = []
@@ -1673,6 +1707,7 @@ class CrossAttnUpBlock2DReferenceOnly(CrossAttnUpBlock2D):
style_fidelity * hidden_states_c
+ (1.0 - style_fidelity) * hidden_states_uc
)
hidden_states *= self.ref_data.strength
if MODE == "read":
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 = ""