diff --git a/ella_example_workflow.json b/ella_example_workflow.json index f647b86..5de711a 100644 --- a/ella_example_workflow.json +++ b/ella_example_workflow.json @@ -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" ] diff --git a/model.py b/model.py index d5ce69f..f8235da 100644 --- a/model.py +++ b/model.py @@ -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) diff --git a/nodes.py b/nodes.py index e8b5c66..5dab88b 100644 --- a/nodes.py +++ b/nodes.py @@ -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) @@ -331,20 +259,64 @@ 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" }