This commit is contained in:
kijai
2024-04-09 21:17:11 +03:00
parent 835020fa7b
commit 1063c40530
3 changed files with 230 additions and 229 deletions
+155 -115
View File
@@ -1,85 +1,13 @@
{
"last_node_id": 6,
"last_link_id": 10,
"last_node_id": 35,
"last_link_id": 34,
"nodes": [
{
"id": 4,
"type": "PreviewImage",
"pos": [
1401,
136
],
"size": {
"0": 593,
"1": 648
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 10
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 5,
"type": "ella_model_loader",
"pos": [
688,
144
],
"size": {
"0": 210,
"1": 66
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 5
},
{
"name": "clip",
"type": "CLIP",
"link": 6,
"slot_index": 1
},
{
"name": "vae",
"type": "VAE",
"link": 7
}
],
"outputs": [
{
"name": "ella_model",
"type": "ELLAMODEL",
"links": [
9
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ella_model_loader"
}
},
{
"id": 3,
"id": 29,
"type": "CheckpointLoaderSimple",
"pos": [
310,
142
289,
315
],
"size": {
"0": 315,
@@ -93,7 +21,7 @@
"name": "MODEL",
"type": "MODEL",
"links": [
5
20
],
"shape": 3,
"slot_index": 0
@@ -102,7 +30,7 @@
"name": "CLIP",
"type": "CLIP",
"links": [
6
21
],
"shape": 3
},
@@ -110,7 +38,7 @@
"name": "VAE",
"type": "VAE",
"links": [
7
22
],
"shape": 3,
"slot_index": 2
@@ -120,28 +48,93 @@
"Node name for S&R": "CheckpointLoaderSimple"
},
"widgets_values": [
"1_5/v1-5-pruned.ckpt"
"1_5\\photon_v1.safetensors"
]
},
{
"id": 6,
"id": 30,
"type": "PreviewImage",
"pos": [
1316,
307
],
"size": {
"0": 590.6172485351562,
"1": 614.5595092773438
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 34
}
],
"properties": {
"Node name for S&R": "PreviewImage"
}
},
{
"id": 34,
"type": "ella_t5_embeds",
"pos": [
512,
484
],
"size": {
"0": 400,
"1": 200
},
"flags": {},
"order": 1,
"mode": 0,
"outputs": [
{
"name": "ella_embeds",
"type": "ELLAEMBEDS",
"links": [
33
],
"shape": 3
}
],
"properties": {
"Node name for S&R": "ella_t5_embeds"
},
"widgets_values": [
"A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",
4,
128,
false
]
},
{
"id": 35,
"type": "ella_sampler",
"pos": [
947,
142
],
"size": [
415,
487
962,
317
],
"size": {
"0": 315,
"1": 222
},
"flags": {},
"order": 2,
"order": 3,
"mode": 0,
"inputs": [
{
"name": "ella_model",
"type": "ELLAMODEL",
"link": 9
"link": 32
},
{
"name": "ella_embeds",
"type": "ELLAEMBEDS",
"link": 33,
"slot_index": 1
}
],
"outputs": [
@@ -149,72 +142,119 @@
"name": "images",
"type": "IMAGE",
"links": [
10
34
],
"shape": 3,
"slot_index": 0
},
{
"name": "last_image",
"type": "IMAGE",
"links": null,
"shape": 3
}
],
"properties": {
"Node name for S&R": "ella_sampler"
},
"widgets_values": [
"A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.\n",
512,
512,
1,
25,
10,
933038223352312,
915981713542918,
"randomize",
"DDPMScheduler"
]
},
{
"id": 27,
"type": "ella_model_loader",
"pos": [
680,
317
],
"size": {
"0": 210,
"1": 66
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "MODEL",
"link": 20
},
{
"name": "clip",
"type": "CLIP",
"link": 21,
"slot_index": 1
},
{
"name": "vae",
"type": "VAE",
"link": 22
}
],
"outputs": [
{
"name": "ella_model",
"type": "ELLAMODEL",
"links": [
32
],
"shape": 3,
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ella_model_loader"
}
}
],
"links": [
[
5,
3,
20,
29,
0,
5,
27,
0,
"MODEL"
],
[
6,
3,
21,
29,
1,
5,
27,
1,
"CLIP"
],
[
7,
3,
22,
29,
2,
5,
27,
2,
"VAE"
],
[
9,
5,
32,
27,
0,
6,
35,
0,
"ELLAMODEL"
],
[
10,
6,
33,
34,
0,
4,
35,
1,
"ELLAEMBEDS"
],
[
34,
35,
0,
30,
0,
"IMAGE"
]
+1 -1
View File
@@ -128,7 +128,7 @@ class PerceiverResampler(nn.Module):
class T5TextEmbedder(nn.Module):
def __init__(self, pretrained_path="google/flan-t5-xl", max_length=None):
def __init__(self, pretrained_path="ybelkada/flan-t5-xl-sharded-bf16", max_length=None):
super().__init__()
self.model = T5EncoderModel.from_pretrained(pretrained_path)
self.tokenizer = T5Tokenizer.from_pretrained(pretrained_path)
+72 -111
View File
@@ -73,72 +73,6 @@ class ELLAProxyUNet(torch.nn.Module):
encoder_attention_mask=encoder_attention_mask,
return_dict=return_dict,
)
def generate_image_with_flexible_max_length(
pipe, t5_encoder, prompt, fixed_negative=False, output_type="pt", **pipe_kwargs
):
device = pipe.device
dtype = pipe.dtype
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
negative_prompt_embeds = t5_encoder(
[""] * batch_size, max_length=128 if fixed_negative else None
).to(device, dtype)
# diffusers pipeline concatenate `prompt_embeds` too early...
# https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913
pipe.unet.flexible_max_length_workaround = [
negative_prompt_embeds.size(1)
] * batch_size + [prompt_embeds.size(1)] * batch_size
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
b, _, d = prompt_embeds.shape
prompt_embeds = torch.cat(
[
prompt_embeds,
torch.zeros(
(b, max_length - prompt_embeds.size(1), d), device=device, dtype=dtype
),
],
dim=1,
)
negative_prompt_embeds = torch.cat(
[
negative_prompt_embeds,
torch.zeros(
(b, max_length - negative_prompt_embeds.size(1), d),
device=device,
dtype=dtype,
),
],
dim=1,
)
images = pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
**pipe_kwargs,
output_type=output_type,
).images
pipe.unet.flexible_max_length_workaround = None
return images
def load_ella(filename, device, dtype):
ella = ELLA()
safetensors.torch.load_model(ella, filename, strict=True)
ella.to(device, dtype=dtype)
return ella
def load_ella_for_pipe(pipe, ella):
pipe.unet = ELLAProxyUNet(ella, pipe.unet)
def offload_ella_for_pipe(pipe):
pipe.unet = pipe.unet.unet
def generate_image_with_fixed_max_length(
pipe, t5_encoder, prompt, output_type="pt", **pipe_kwargs
@@ -176,6 +110,7 @@ class ella_model_loader:
def loadmodel(self, model, clip, vae):
mm.soft_empty_cache()
dtype = mm.unet_dtype()
vae_dtype = mm.vae_dtype()
device = mm.get_torch_device()
custom_config = {
@@ -231,24 +166,21 @@ class ella_model_loader:
'beta_schedule': "linear",
'steps_offset': 1
}
# 4. tokenizer
tokenizer_path = os.path.join(script_directory, "configs/tokenizer")
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
scheduler=DPMSolverMultistepScheduler(**scheduler_config)
pbar.update(1)
del sd
print("loading ELLA")
ella_path = os.path.join(script_directory, 'checkpoints', 'ella-sd1.5-tsc-t5xl.safetensors')
ella = ELLA()
safetensors.torch.load_model(ella, ella_path, strict=True)
ella.to(device, dtype=dtype)
unet = unet.to(device)
ella_unet = ELLAProxyUNet(ella, unet)
pbar.update(1)
print("loading tokenizer")
tokenizer_path = os.path.join(script_directory, "configs/tokenizer")
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path)
print("creating pipeline")
pipe = StableDiffusionPipeline(
self.pipe = StableDiffusionPipeline(
unet=unet,
vae=vae,
text_encoder=text_encoder,
@@ -261,12 +193,10 @@ class ella_model_loader:
)
print("pipeline created")
pbar.update(1)
pipe.unet = ella_unet
t5_encoder = T5TextEmbedder().to(pipe.device, dtype=dtype)
self.pipe.unet = ella_unet
ella_model = {
'pipe': pipe,
'ella': ella,
't5_encoder': t5_encoder
'pipe': self.pipe,
}
return (ella_model,)
@@ -276,10 +206,9 @@ class ella_sampler:
def INPUT_TYPES(s):
return {"required": {
"ella_model": ("ELLAMODEL",),
"prompt": ("STRING", {"multiline": True, "default": "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",}),
"ella_embeds": ("ELLAEMBEDS",),
"width": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
"height": ("INT", {"default": 512, "min": 64, "max": 2048, "step": 64}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256, "step": 1}),
"steps": ("INT", {"default": 25, "min": 1, "max": 200, "step": 1}),
"guidance_scale": ("FLOAT", {"default": 10.0, "min": 0.0, "max": 20.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
@@ -291,7 +220,7 @@ class ella_sampler:
'PNDMScheduler',
'DEISMultistepScheduler'
], {
"default": 'DDIMScheduler'
"default": 'DDPMScheduler'
}),
},
}
@@ -299,14 +228,13 @@ class ella_sampler:
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("images",)
FUNCTION = "process"
CATEGORY = "champWrapper"
CATEGORY = "ELLA-Wrapper"
def process(self, prompt, batch_size, width, height, steps, guidance_scale, seed, ella_model, scheduler):
def process(self, ella_embeds, width, height, steps, guidance_scale, seed, ella_model, scheduler):
device = mm.get_torch_device()
mm.unload_all_models()
mm.soft_empty_cache()
dtype = mm.unet_dtype()
t5_encoder=ella_model['t5_encoder']
pipe=ella_model['pipe']
pipe.to(device, dtype=dtype)
@@ -332,19 +260,63 @@ class ella_sampler:
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
# diffusers pipeline concatenate `prompt_embeds` too early...
# https://github.com/huggingface/diffusers/blob/b6d7e31d10df675d86c6fe7838044712c6dca4e9/src/diffusers/pipelines/stable_diffusion/pipeline_stable_diffusion.py#L913
pipe.unet.flexible_max_length_workaround = [ella_embeds["negative_prompt_embeds"].size(1)] * ella_embeds["batch_size"] + [ella_embeds["prompt_embeds"].size(1)] * ella_embeds["batch_size"]
images = pipe(
prompt_embeds=ella_embeds["prompt_embeds"],
negative_prompt_embeds=ella_embeds["negative_prompt_embeds"],
guidance_scale=guidance_scale,
num_inference_steps=steps,
height=height,
width=width,
generator=[
torch.Generator(device=device).manual_seed(seed + i)
for i in range(ella_embeds["batch_size"])
],
output_type="np.array",
).images
image_out = torch.from_numpy(images).cpu().float()
return (image_out,)
class ella_t5_embeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"prompt": ("STRING", {"multiline": True, "default": "A vivid red book with a smooth, matte cover lies next to a glossy yellow vase. The vase, with a slightly curved silhouette, stands on a dark wood table with a noticeable grain pattern. The book appears slightly worn at the edges, suggesting frequent use, while the vase holds a fresh array of multicolored wildflowers.",}),
"batch_size": ("INT", {"default": 1, "min": 1, "max": 256, "step": 1}),
"max_length": ("INT", {"default": 128, "min": 1, "max": 256, "step": 1}),
"fixed_negative": ("BOOLEAN", {"default": False}),
},
}
RETURN_TYPES = ("ELLAEMBEDS",)
RETURN_NAMES = ("ella_embeds",)
FUNCTION = "process"
CATEGORY = "ELLA-Wrapper"
def process(self, prompt, batch_size, max_length, fixed_negative):
device = mm.get_torch_device()
mm.unload_all_models()
mm.soft_empty_cache()
dtype = mm.unet_dtype()
t5_encoder = T5TextEmbedder().to(device, dtype=dtype)
autocast_condition = (dtype != torch.float32) and not mm.is_device_mps(device)
with torch.autocast(mm.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext():
print("generating embeds")
prompt = [prompt] * batch_size
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
#batch_size = len(prompt)
fixed_negative = False
prompt_embeds = t5_encoder(prompt, max_length=None).to(device, dtype)
negative_prompt_embeds = t5_encoder(
[""] * batch_size, max_length=128 if fixed_negative else None
[""] * batch_size, max_length=max_length if fixed_negative else None
).to(device, dtype)
pipe.unet.flexible_max_length_workaround = [
negative_prompt_embeds.size(1)
] * batch_size + [prompt_embeds.size(1)] * batch_size
max_length = max([prompt_embeds.size(1), negative_prompt_embeds.size(1)])
b, _, d = prompt_embeds.shape
prompt_embeds = torch.cat(
@@ -367,31 +339,20 @@ class ella_sampler:
],
dim=1,
)
images = pipe(
prompt_embeds=prompt_embeds,
negative_prompt_embeds=negative_prompt_embeds,
guidance_scale=guidance_scale,
num_inference_steps=steps,
height=height,
width=width,
generator=[
torch.Generator(device=device).manual_seed(seed + i)
for i in range(batch_size)
],
output_type="np.array",
).images
tensor = torch.from_numpy(images).cpu().float()
return (tensor,)
embeds = {
"prompt_embeds": prompt_embeds,
"negative_prompt_embeds": negative_prompt_embeds,
"batch_size": batch_size
}
return (embeds,)
NODE_CLASS_MAPPINGS = {
"ella_model_loader": ella_model_loader,
"ella_sampler": ella_sampler,
"ella_t5_embeds": ella_t5_embeds
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ella_model_loader": "ELLA Model Loader",
"ella_sampler": "ELLA Sampler",
"ella_t5_embeds": "ELLA T5 Embeds"
}