Merge pull request #22 from JettHu/main
feat: support multi batch, lora tirgger words, prompt weight with clip
This commit is contained in:
@@ -11,6 +11,7 @@
|
||||
|
||||
## :star2: Changelog
|
||||
|
||||
- **[2024.4.22]** Fix unstable quality of image while multi-batch. Add CLIP concat (support lora trigger words now).
|
||||
- **[2024.4.19]** Documenting nodes
|
||||
- **[2024.4.19]** Initial repo
|
||||
|
||||
@@ -20,10 +21,16 @@ The [examples directory](./examples/) has workflow examples. You can directly lo
|
||||
|
||||

|
||||
|
||||
:tada: It works with controlnet! And [EMMA](https://github.com/TencentQQGYLab/ELLA/issues/15) is working in progress.
|
||||
:tada: It works with controlnet!
|
||||
|
||||

|
||||
|
||||
:tada: It works with **lora trigger words** by concat CLIP CONDITION!
|
||||
|
||||

|
||||
|
||||
And [EMMA](https://github.com/TencentQQGYLab/ELLA/issues/15) is working in progress.
|
||||
|
||||
## :green_book: Install
|
||||
|
||||
Download or git clone this repository inside ComfyUI/custom_nodes/ directory. `ComfyUI-ELLA` requires the latest version of ComfyUI. If something doesn't work be sure to upgrade.
|
||||
@@ -72,6 +79,10 @@ ComfyUI/models/ella_encoder/
|
||||
|
||||
[Nodes reference](./NODES.md)
|
||||
|
||||
## :mag: Common promblem
|
||||
|
||||
- XXX not implemented for 'Half'. See [issue #12](https://github.com/TencentQQGYLab/ComfyUI-ELLA/issues/12#issuecomment-2067994702)
|
||||
|
||||
## :memo: TODO
|
||||
|
||||
- [ ] Support prompt weighting
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.6 MiB |
@@ -1,8 +1,10 @@
|
||||
import os
|
||||
from typing import Dict
|
||||
|
||||
import folder_paths
|
||||
import torch
|
||||
from comfy import model_management
|
||||
from comfy.conds import CONDCrossAttn
|
||||
from safetensors.torch import load_model
|
||||
|
||||
from .model import ELLA, T5TextEmbedder
|
||||
@@ -10,6 +12,8 @@ from .model import ELLA, T5TextEmbedder
|
||||
ELLA_EMBEDS_TYPE = "ELLA_EMBEDS"
|
||||
ELLA_EMBEDS_PREFIX = "ella_"
|
||||
ELLA_EMBEDS_PREFIX_LEN = len(ELLA_EMBEDS_PREFIX)
|
||||
APPLY_MODE_ELLA_ONLY = "ELLA ONLY"
|
||||
APPLY_MODE_ELLA_AND_CLIP = "ELLA + CLIP"
|
||||
|
||||
# set the models directory
|
||||
if "ella" not in folder_paths.folder_names_and_paths:
|
||||
@@ -26,18 +30,34 @@ folder_paths.folder_names_and_paths["ella_encoder"] = (current_paths, folder_pat
|
||||
|
||||
|
||||
class EllaProxyUNet:
|
||||
def __init__(self, ella, model_sampling, positive, negative) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
ella,
|
||||
model_sampling,
|
||||
positive,
|
||||
negative,
|
||||
mode=APPLY_MODE_ELLA_ONLY,
|
||||
sigma_start=99999999,
|
||||
sigma_end=0,
|
||||
**kwargs,
|
||||
) -> None:
|
||||
self.ella = ella
|
||||
self.model_sampling = model_sampling
|
||||
self.sigma_start = sigma_start
|
||||
self.sigma_end = sigma_end
|
||||
self.mode = mode
|
||||
if positive.keys() != negative.keys():
|
||||
raise ValueError("positive and negative embeds types must match")
|
||||
# if mode == APPLY_MODE_ELLA_AND_CLIP and "clip_embeds" not in positive:
|
||||
# raise ValueError(f"'clip_embeds' is required when using '{APPLY_MODE_ELLA_AND_CLIP}' mode")
|
||||
self.embeds = [positive, negative]
|
||||
|
||||
self.dtype = model_management.text_encoder_dtype()
|
||||
self.ella.to(self.dtype)
|
||||
for i in range(len(self.embeds)):
|
||||
for k in self.embeds[i]:
|
||||
self.embeds[i][k].to(device=self.load_device, dtype=self.dtype)
|
||||
self.embeds[i][k].to(dtype=self.dtype)
|
||||
self.embeds[i][k] = CONDCrossAttn(self.embeds[i][k])
|
||||
|
||||
@property
|
||||
def load_device(self):
|
||||
@@ -47,12 +67,25 @@ class EllaProxyUNet:
|
||||
def offload_device(self):
|
||||
return model_management.text_encoder_offload_device()
|
||||
|
||||
def process_cond(self, embeds: Dict[str, CONDCrossAttn], batch_size, **kwargs):
|
||||
return {k: v.process_cond(batch_size, self.load_device, **kwargs).cond for k, v in embeds.items()}
|
||||
|
||||
def prepare_conds(self):
|
||||
self.ella.to(self.load_device)
|
||||
cond = self.ella(torch.Tensor([999]).to(torch.int64), **self.embeds[0])
|
||||
uncond = self.ella(torch.Tensor([999]).to(torch.int64), **self.embeds[1])
|
||||
cond_embeds = self.process_cond(self.embeds[0], 1)
|
||||
cond = self.ella(torch.Tensor([999]).to(torch.int64), **cond_embeds)
|
||||
uncond_embeds = self.process_cond(self.embeds[1], 1)
|
||||
uncond = self.ella(torch.Tensor([999]).to(torch.int64), **uncond_embeds)
|
||||
self.ella.to(self.offload_device)
|
||||
return cond, uncond
|
||||
if self.mode == APPLY_MODE_ELLA_ONLY:
|
||||
return cond, uncond
|
||||
if "clip_embeds" not in cond_embeds or "clip_embeds" not in uncond_embeds:
|
||||
print("warning: 'clip_embeds' is required, fallback to 'ELLA ONLY' mode")
|
||||
return cond, uncond
|
||||
return (
|
||||
torch.concat([cond, cond_embeds["clip_embeds"]], dim=1),
|
||||
torch.concat([uncond, uncond_embeds["clip_embeds"]], dim=1),
|
||||
)
|
||||
|
||||
def __call__(self, apply_model, kwargs: dict):
|
||||
input_x = kwargs["input"]
|
||||
@@ -61,13 +94,22 @@ class EllaProxyUNet:
|
||||
cond_or_uncond = kwargs["cond_or_uncond"] # [0|1]
|
||||
_device = c["c_crossattn"].device
|
||||
|
||||
# TODO: add ella start/end sigma control
|
||||
time_aware_encoder_hidden_states = []
|
||||
self.ella.to(device=self.load_device)
|
||||
for i in cond_or_uncond:
|
||||
cond_embeds = self.process_cond(self.embeds[i], input_x.size(0) // len(cond_or_uncond))
|
||||
h = self.ella(
|
||||
self.model_sampling.timestep(timestep_[i]),
|
||||
**self.embeds[i],
|
||||
self.model_sampling.timestep(timestep_[0]),
|
||||
**cond_embeds,
|
||||
)
|
||||
if self.mode == APPLY_MODE_ELLA_ONLY:
|
||||
time_aware_encoder_hidden_states.append(h)
|
||||
continue
|
||||
if "clip_embeds" not in cond_embeds:
|
||||
time_aware_encoder_hidden_states.append(h)
|
||||
continue
|
||||
h = torch.concat([h, cond_embeds["clip_embeds"]], dim=1)
|
||||
time_aware_encoder_hidden_states.append(h)
|
||||
self.ella.to(self.offload_device)
|
||||
|
||||
@@ -83,7 +125,7 @@ class EllaProxyUNet:
|
||||
"""
|
||||
|
||||
|
||||
class EllaApply:
|
||||
class EllaAdvancedApply:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
@@ -92,7 +134,12 @@ class EllaApply:
|
||||
"ella": ("ELLA",),
|
||||
"positive": (ELLA_EMBEDS_TYPE,),
|
||||
"negative": (ELLA_EMBEDS_TYPE,),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"mode": ([APPLY_MODE_ELLA_AND_CLIP, APPLY_MODE_ELLA_ONLY],),
|
||||
"start_at": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
"end_at": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_NAMES = ("model", "positive", "negative")
|
||||
@@ -100,9 +147,20 @@ class EllaApply:
|
||||
FUNCTION = "apply"
|
||||
CATEGORY = "ella/apply"
|
||||
|
||||
def apply(self, model, ella, positive, negative):
|
||||
def apply(
|
||||
self,
|
||||
model,
|
||||
ella,
|
||||
positive,
|
||||
negative,
|
||||
mode=APPLY_MODE_ELLA_AND_CLIP,
|
||||
start_at=0.0,
|
||||
end_at=1.0,
|
||||
):
|
||||
model_clone = model.clone()
|
||||
model_sampling = model_clone.get_model_object("model_sampling")
|
||||
sigma_start = model_clone.get_model_object("model_sampling").percent_to_sigma(start_at)
|
||||
sigma_end = model_clone.get_model_object("model_sampling").percent_to_sigma(end_at)
|
||||
|
||||
ella_proxy = EllaProxyUNet(
|
||||
ella=ella,
|
||||
@@ -113,6 +171,9 @@ class EllaApply:
|
||||
negative={
|
||||
k[ELLA_EMBEDS_PREFIX_LEN:]: v.clone() for k, v in negative.items() if k.startswith(ELLA_EMBEDS_PREFIX)
|
||||
},
|
||||
mode=mode,
|
||||
sigma_start=sigma_start,
|
||||
sigma_end=sigma_end,
|
||||
)
|
||||
|
||||
model_clone.set_model_unet_function_wrapper(ella_proxy)
|
||||
@@ -124,6 +185,19 @@ class EllaApply:
|
||||
return (model_clone, [cond], [uncond])
|
||||
|
||||
|
||||
class EllaApply(EllaAdvancedApply):
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"model": ("MODEL",),
|
||||
"ella": ("ELLA",),
|
||||
"positive": (ELLA_EMBEDS_TYPE,),
|
||||
"negative": (ELLA_EMBEDS_TYPE,),
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
"""
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
Encoders
|
||||
@@ -274,6 +348,30 @@ class EllaCombineEmbeds:
|
||||
return ({**embeds, **embeds_add},)
|
||||
|
||||
|
||||
class ConcatConditionEllaEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"required": {
|
||||
"cond": ("CONDITIONING",),
|
||||
"embeds": (ELLA_EMBEDS_TYPE,),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = (ELLA_EMBEDS_TYPE,)
|
||||
FUNCTION = "combine"
|
||||
|
||||
CATEGORY = "ella/helper"
|
||||
|
||||
def combine(self, cond, embeds):
|
||||
# only use batch 0
|
||||
# CONDITIONING: [[cond, {"pooled_output": pooled}]]
|
||||
clip_key = f"{ELLA_EMBEDS_PREFIX}clip_embeds"
|
||||
if clip_key in embeds:
|
||||
print("warning: there is already a clip embeds, the previous condition will be overwritten")
|
||||
return ({f"{ELLA_EMBEDS_PREFIX}clip_embeds": cond[0][0], **cond[0][1], **embeds},)
|
||||
|
||||
|
||||
"""
|
||||
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
|
||||
Register
|
||||
@@ -282,6 +380,7 @@ class EllaCombineEmbeds:
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
# Main Apply Nodes
|
||||
"EllaApply": EllaApply,
|
||||
# "EllaAdvancedApply": EllaAdvancedApply,
|
||||
"T5TextEncode #ELLA": T5TextEncode,
|
||||
# Loaders
|
||||
"ELLALoader": ELLALoader,
|
||||
@@ -289,11 +388,13 @@ NODE_CLASS_MAPPINGS = {
|
||||
# Helpers
|
||||
"EllaCombineEmbeds": EllaCombineEmbeds,
|
||||
"ConditionToEllaEmbeds": ConditionToEllaEmbeds,
|
||||
"ConcatConditionEllaEmbeds": ConcatConditionEllaEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Main Apply Nodes
|
||||
"EllaApply": "Apply ELLA",
|
||||
# "EllaAdvancedApply": "Apply ELLA Advanced",
|
||||
"T5TextEncode #ELLA": "T5 Text Encode #ELLA",
|
||||
# Loaders
|
||||
"ELLALoader": "Load ELLA Model",
|
||||
@@ -301,4 +402,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
# Helpers
|
||||
"EllaCombineEmbeds": "ELLA Combine Embeds",
|
||||
"ConditionToEllaEmbeds": "Convert Condition to ELLA Embeds",
|
||||
"ConcatConditionEllaEmbeds": "Concat Condition & ELLA Embeds",
|
||||
}
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 524 KiB |
@@ -1,3 +1,4 @@
|
||||
torch
|
||||
safetensors
|
||||
transformers
|
||||
sentencepiece
|
||||
|
||||
Reference in New Issue
Block a user