From 99acdc520d771ac21396daa762664b3c1fe8d21d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 12 Oct 2024 14:31:32 +0300 Subject: [PATCH] text2vid multiprompt batches --- ...text2video_multipleprompts_example_01.json | 457 ++++++++++++++++++ nodes.py | 88 +++- .../pyramid_dit_for_video_gen_pipeline.py | 7 +- 3 files changed, 541 insertions(+), 11 deletions(-) create mode 100644 examples/pyramidflow_text2video_multipleprompts_example_01.json diff --git a/examples/pyramidflow_text2video_multipleprompts_example_01.json b/examples/pyramidflow_text2video_multipleprompts_example_01.json new file mode 100644 index 0000000..9f6e9b8 --- /dev/null +++ b/examples/pyramidflow_text2video_multipleprompts_example_01.json @@ -0,0 +1,457 @@ +{ + "last_node_id": 29, + "last_link_id": 43, + "nodes": [ + { + "id": 9, + "type": "PyramidFlowSampler", + "pos": { + "0": 1059, + "1": 497 + }, + "size": { + "0": 411.5168151855469, + "1": 314 + }, + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 7 + }, + { + "name": "prompt_embeds", + "type": "PYRAMIDFLOWPROMPT", + "link": 43 + }, + { + "name": "input_latent", + "type": "LATENT", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "links": [ + 8 + ] + }, + { + "name": "samples", + "type": "LATENT", + "links": [ + 9 + ], + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowSampler" + }, + "widgets_values": [ + 1280, + 768, + "20, 20, 20", + "10, 10, 10", + 16, + 7, + 5, + 44664248661394, + "fixed", + "" + ] + }, + { + "id": 8, + "type": "PyramidFlowVAEDecode", + "pos": { + "0": 1161, + "1": 873 + }, + "size": { + "0": 315, + "1": 102 + }, + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 8 + }, + { + "name": "samples", + "type": "LATENT", + "link": 9 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 39 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowVAEDecode" + }, + "widgets_values": [ + 256, + 2 + ] + }, + { + "id": 28, + "type": "GetImageSizeAndCount", + "pos": { + "0": 1180, + "1": 1118 + }, + "size": { + "0": 277.20001220703125, + "1": 86 + }, + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 39 + } + ], + "outputs": [ + { + "name": "image", + "type": "IMAGE", + "links": [ + 40 + ], + "slot_index": 0 + }, + { + "name": "1280 width", + "type": "INT", + "links": null + }, + { + "name": "768 height", + "type": "INT", + "links": null + }, + { + "name": "242 count", + "type": "INT", + "links": null + } + ], + "properties": { + "Node name for S&R": "GetImageSizeAndCount" + }, + "widgets_values": [] + }, + { + "id": 5, + "type": "DownloadAndLoadPyramidFlowModel", + "pos": { + "0": 576, + "1": 496 + }, + "size": { + "0": 385.7839050292969, + "1": 202 + }, + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "pyramidflow_model", + "type": "PYRAMIDFLOWMODEL", + "links": [ + 7, + 30, + 41 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadPyramidFlowModel" + }, + "widgets_values": [ + "rain1011/pyramid-flow-sd3", + "diffusion_transformer_768p", + "bf16", + "bf16", + "bf16", + false, + false + ] + }, + { + "id": 29, + "type": "PyramidFlowTextEncode", + "pos": { + "0": 570, + "1": 1043 + }, + "size": { + "0": 434.50982666015625, + "1": 227.74803161621094 + }, + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 41 + }, + { + "name": "prev_prompt", + "type": "PYRAMIDFLOWPROMPT", + "link": 42, + "shape": 7 + } + ], + "outputs": [ + { + "name": "prompt_embeds", + "type": "PYRAMIDFLOWPROMPT", + "links": [ + 43 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowTextEncode" + }, + "widgets_values": [ + "A massive explosion on the surface of the earth, hyper quality, Ultra HD, 8K", + "cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", + false + ] + }, + { + "id": 22, + "type": "PyramidFlowTextEncode", + "pos": { + "0": 567, + "1": 757 + }, + "size": { + "0": 434.50982666015625, + "1": 227.74803161621094 + }, + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 30 + }, + { + "name": "prev_prompt", + "type": "PYRAMIDFLOWPROMPT", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "prompt_embeds", + "type": "PYRAMIDFLOWPROMPT", + "links": [ + 42 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowTextEncode" + }, + "widgets_values": [ + "A campfire burning with flames and embers, gradually increasing in size and intensity before dying down towards the end, hyper quality, Ultra HD, 8K", + "cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", + true + ] + }, + { + "id": 14, + "type": "VHS_VideoCombine", + "pos": { + "0": 1534, + "1": 490 + }, + "size": [ + 1700, + 1332 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 40 + }, + { + "name": "audio", + "type": "AUDIO", + "link": null, + "shape": 7 + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null, + "shape": 7 + }, + { + "name": "vae", + "type": "VAE", + "link": null, + "shape": 7 + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "PyramidFlow", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": true, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "PyramidFlow_00038.mp4", + "subfolder": "", + "type": "output", + "format": "video/h264-mp4", + "frame_rate": 16 + }, + "muted": false + } + } + } + ], + "links": [ + [ + 7, + 5, + 0, + 9, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 8, + 9, + 0, + 8, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 9, + 9, + 1, + 8, + 1, + "LATENT" + ], + [ + 30, + 5, + 0, + 22, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 39, + 8, + 0, + 28, + 0, + "IMAGE" + ], + [ + 40, + 28, + 0, + 14, + 0, + "IMAGE" + ], + [ + 41, + 5, + 0, + 29, + 0, + "PYRAMIDFLOWMODEL" + ], + [ + 42, + 22, + 0, + 29, + 1, + "PYRAMIDFLOWPROMPT" + ], + [ + 43, + 29, + 0, + 9, + 1, + "PYRAMIDFLOWPROMPT" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.6934334949442883, + "offset": [ + -378.21925980506256, + -283.47815759899163 + ] + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/nodes.py b/nodes.py index adcad45..cfdd0e6 100644 --- a/nodes.py +++ b/nodes.py @@ -36,7 +36,6 @@ class DownloadAndLoadPyramidFlowModel: "model_dtype": (["fp8_e4m3fn","fp8_e5m2","fp16", "fp32", "bf16"],{"default": "bf16", }), "text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), "vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), - "use_flash_attn": ("BOOLEAN", {"default": False}), "fp8_fastmode": ("BOOLEAN",{"default": False, "tooltip": "fastmode is only for latest nvidia GPUs"}), #"compile": (["disabled","onediff","torch"], {"tooltip": "compile the model for faster inference, these are advanced options only available on Linux, see readme for more info"}), } @@ -47,7 +46,7 @@ class DownloadAndLoadPyramidFlowModel: FUNCTION = "loadmodel" CATEGORY = "PyramidFlowWrapper" - def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, fp8_fastmode, use_flash_attn=False): + def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, fp8_fastmode): device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -87,7 +86,6 @@ class DownloadAndLoadPyramidFlowModel: text_encoder_dtype, vae_dtype, model_variant=variant, - use_flash_attn=use_flash_attn, fp8_fastmode=fp8_fastmode, ) @@ -241,9 +239,9 @@ class PyramidFlowTextEncode: "keep_model_loaded": ("BOOLEAN", {"default": False}), }, - # "optional": { - # "samples": ("LATENT", ), - # } + "optional": { + "prev_prompt": ("PYRAMIDFLOWPROMPT", ), + } } RETURN_TYPES = ("PYRAMIDFLOWPROMPT", ) @@ -251,7 +249,7 @@ class PyramidFlowTextEncode: FUNCTION = "sample" CATEGORY = "PyramidFlowWrapper" - def sample(self, model, positive_prompt, negative_prompt, keep_model_loaded): + def sample(self, model, positive_prompt, negative_prompt, keep_model_loaded, prev_prompt=None): mm.soft_empty_cache() device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -268,6 +266,15 @@ class PyramidFlowTextEncode: if not keep_model_loaded: text_encoder.to(offload_device) + if prev_prompt is not None: + prompt_embeds = torch.cat((prev_prompt["prompt_embeds"], prompt_embeds), dim=0) + prompt_attention_mask = torch.cat((prev_prompt["attention_mask"], prompt_attention_mask), dim=0) + pooled_prompt_embeds = torch.cat((prev_prompt["pooled_embeds"], pooled_prompt_embeds), dim=0) + + negative_prompt_embeds = torch.cat((prev_prompt["negative_prompt_embeds"], negative_prompt_embeds), dim=0) + negative_prompt_attention_mask = torch.cat((prev_prompt["negative_attention_mask"], negative_prompt_attention_mask), dim=0) + pooled_negative_prompt_embeds = torch.cat((prev_prompt["negative_pooled_embeds"], pooled_negative_prompt_embeds), dim=0) + embeds = { "prompt_embeds": prompt_embeds, "attention_mask": prompt_attention_mask, @@ -278,7 +285,70 @@ class PyramidFlowTextEncode: } return (embeds,) - + +#not functional yet, todo: figure out why the results are bad with it +class PyramidFlowTextEncodeComfy: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "clip": ("CLIP",), + "positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ), + "negative_prompt": ("STRING", {"default": "", "multiline": True} ), + } + } + + RETURN_TYPES = ("PYRAMIDFLOWPROMPT",) + RETURN_NAMES = ("prompt_embeds",) + FUNCTION = "process" + CATEGORY = "CogVideoWrapper" + + def process(self, clip, positive_prompt, negative_prompt): + device = mm.get_torch_device() + offload_device = mm.unet_offload_device() + clip.cond_stage_model.reset_clip_options() + clip.tokenizer.t5xxl.pad_to_max_length = True + clip.tokenizer.t5xxl.truncation = True + clip.tokenizer.t5xxl.max_length = 128 + clip.cond_stage_model.t5xxl.return_attention_masks = True + clip.cond_stage_model.t5_attention_mask = True + + + clip.cond_stage_model.t5xxl.to(device) + tokens = clip.tokenize(positive_prompt.lower().strip(), return_word_ids=True) + + prompt_embeds, pooled_prompt_embeds, prompt_attention_mask = clip.cond_stage_model.encode_token_weights(tokens) + tokens = clip.tokenize(negative_prompt.lower().strip(), return_word_ids=True) + negative_prompt_embeds, pooled_negative_prompt_embeds, negative_prompt_attention_mask = clip.cond_stage_model.encode_token_weights(tokens) + clip.cond_stage_model.t5xxl.to(offload_device) + + max_length = prompt_attention_mask["attention_mask"].shape[1] + prompt_embeds = prompt_embeds[:, :max_length, :] + + print(prompt_embeds.shape) + print(prompt_attention_mask["attention_mask"].shape) + + # If the sequence length is less than max_length, pad the embeddings + if prompt_embeds.shape[1] < max_length: + padding = torch.zeros((prompt_embeds.shape[0], max_length - prompt_embeds.shape[1], prompt_embeds.shape[2]), device=prompt_embeds.device) + prompt_embeds = torch.cat((prompt_embeds, padding), dim=1) + + max_length = negative_prompt_attention_mask["attention_mask"].shape[1] + negative_prompt_embeds = negative_prompt_embeds[:, :max_length, :] + + if negative_prompt_embeds.shape[1] < max_length: + padding = torch.zeros((negative_prompt_embeds.shape[0], max_length - negative_prompt_embeds.shape[1], negative_prompt_embeds.shape[2]), device=negative_prompt_embeds.device) + negative_prompt_embeds = torch.cat((negative_prompt_embeds, padding), dim=1) + + embeds = { + "prompt_embeds": prompt_embeds.to(device), + "attention_mask": prompt_attention_mask["attention_mask"].to(device), + "pooled_embeds": pooled_prompt_embeds.to(device), + "negative_prompt_embeds": negative_prompt_embeds.to(device), + "negative_attention_mask": negative_prompt_attention_mask["attention_mask"].to(device), + "negative_pooled_embeds": pooled_negative_prompt_embeds.to(device) + } + + return (embeds, ) class PyramidFlowVAEEncode: @classmethod def INPUT_TYPES(s): @@ -384,6 +454,7 @@ NODE_CLASS_MAPPINGS = { "PyramidFlowVAEDecode": PyramidFlowVAEDecode, "PyramidFlowTextEncode": PyramidFlowTextEncode, "PyramidFlowVAEEncode": PyramidFlowVAEEncode, + #"PyramidFlowTextEncodeComfy": PyramidFlowTextEncodeComfy, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -392,4 +463,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "PyramidFlowVAEDecode" : "PyramidFlow VAE Decode", "PyramidFlowTextEncode": "PyramidFlow Text Encode", "PyramidFlowVAEEncode": "PyramidFlow VAE Encode", + #"PyramidFlowTextEncodeComfy": "PyramidFlow Text Encode Comfy", } diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index 5d88121..f8ada58 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -311,7 +311,7 @@ class PyramidDiTForVideoGeneration: dtype = self.dtype assert temp % self.frame_per_unit == 0, "The frames should be divided by frame_per unit" - batch_size = 1 + batch_size = prompt_embeds_dict['prompt_embeds'].shape[0] # if isinstance(prompt, str): # batch_size = 1 # prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics @@ -391,7 +391,8 @@ class PyramidDiTForVideoGeneration: #input_image_tensor = image_transform(input_image).unsqueeze(0).unsqueeze(2) # [b c 1 h w] input_image_latent = input_image_latent.to(dtype).to(device) - generated_latents_list = [input_image_latent] # The generated results + #generated_latents_list = [input_image_latent] # The generated results + generated_latents_list = list(torch.unbind(input_image_latent, dim=0)) #last_generated_latents = input_image_latent self.dit.to(device) @@ -507,7 +508,7 @@ class PyramidDiTForVideoGeneration: # negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device) # self.text_encoder.to('cpu') - batch_size=1 + batch_size = prompt_embeds_dict['prompt_embeds'].shape[0] if use_linear_guidance: max_guidance_scale = guidance_scale