Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ab4a4fb38c | ||
|
|
92b0b83839 | ||
|
|
17a0274d77 | ||
|
|
3c121b126c | ||
|
|
4f6bb64067 | ||
|
|
4b2aac5426 | ||
|
|
59326f53d1 | ||
|
|
aa0211ee28 | ||
|
|
55636490a3 | ||
|
|
8698c62af7 | ||
|
|
d3f732236a | ||
|
|
5f94a0cf46 | ||
|
|
371e006976 | ||
|
|
49c360978a | ||
|
|
f3c696d5db | ||
|
|
d3edc4a8f4 | ||
|
|
5be339b5a0 | ||
|
|
0d9da4ebc6 | ||
|
|
83c9d84e3e | ||
|
|
83263b4f96 |
@@ -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 }}
|
||||
@@ -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.
|
||||
|
||||

|
||||
|
||||
## Inpainting with Jannchie's Diffusers Pipeline
|
||||
@@ -58,6 +60,10 @@ A checkpoint for stablediffusion 1.5 is all your need. But for full automation,
|
||||
|
||||

|
||||
|
||||
## QR Code
|
||||
|
||||

|
||||
|
||||
## FAQ
|
||||
|
||||
### Why Diffusers?
|
||||
|
||||
@@ -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",
|
||||
|
||||
|
Before Width: | Height: | Size: 362 KiB After Width: | Height: | Size: 416 KiB |
|
Before Width: | Height: | Size: 2.5 MiB After Width: | Height: | Size: 2.4 MiB |
|
Before Width: | Height: | Size: 1.1 MiB After Width: | Height: | Size: 1.2 MiB |
|
Before Width: | Height: | Size: 792 KiB After Width: | Height: | Size: 806 KiB |
|
After Width: | Height: | Size: 1.7 MiB |
|
Before Width: | Height: | Size: 687 KiB After Width: | Height: | Size: 685 KiB |
@@ -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()
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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 = ""
|
||||