From a137f7a6d2585aa676605425b7f48ce5a0db55ea Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 30 Oct 2024 02:56:32 +0200 Subject: [PATCH] use comfy text encoders for miniflux --- examples/pyramidflow_miniflux_example_01.json | 357 ++++++++++-------- nodes.py | 59 ++- .../flux_modules/modeling_text_encoder.py | 7 + .../pyramid_dit_for_video_gen_pipeline.py | 5 +- 4 files changed, 241 insertions(+), 187 deletions(-) diff --git a/examples/pyramidflow_miniflux_example_01.json b/examples/pyramidflow_miniflux_example_01.json index 909e3ff..84308eb 100644 --- a/examples/pyramidflow_miniflux_example_01.json +++ b/examples/pyramidflow_miniflux_example_01.json @@ -1,134 +1,7 @@ { - "last_node_id": 27, - "last_link_id": 38, + "last_node_id": 39, + "last_link_id": 54, "nodes": [ - { - "id": 8, - "type": "PyramidFlowVAEDecode", - "pos": { - "0": 1161, - "1": 873 - }, - "size": { - "0": 315, - "1": 102 - }, - "flags": {}, - "order": 3, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "PYRAMIDFLOWMODEL", - "link": 8 - }, - { - "name": "samples", - "type": "LATENT", - "link": 9 - } - ], - "outputs": [ - { - "name": "images", - "type": "IMAGE", - "links": [ - 38 - ], - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "PyramidFlowVAEDecode" - }, - "widgets_values": [ - 256, - 2 - ] - }, - { - "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": [ - 31 - ] - } - ], - "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", - false - ] - }, - { - "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 - ], - "slot_index": 0 - } - ], - "properties": { - "Node name for S&R": "DownloadAndLoadPyramidFlowModel" - }, - "widgets_values": [ - "rain1011/pyramid-flow-miniflux", - "diffusion_transformer_384p", - "bf16", - "bf16", - "bf16", - false - ] - }, { "id": 9, "type": "PyramidFlowSampler", @@ -141,7 +14,7 @@ "1": 314 }, "flags": {}, - "order": 2, + "order": 4, "mode": 0, "inputs": [ { @@ -152,7 +25,7 @@ { "name": "prompt_embeds", "type": "PYRAMIDFLOWPROMPT", - "link": 31 + "link": 54 }, { "name": "input_latent", @@ -189,30 +62,74 @@ 16, 9, 5, - 44664248661394, + 44664248661395, "fixed", "" ] }, + { + "id": 8, + "type": "PyramidFlowVAEDecode", + "pos": { + "0": 1161, + "1": 873 + }, + "size": { + "0": 315, + "1": 102 + }, + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "PYRAMIDFLOWMODEL", + "link": 8 + }, + { + "name": "samples", + "type": "LATENT", + "link": 9 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 53 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowVAEDecode" + }, + "widgets_values": [ + 256, + 2 + ] + }, { "id": 14, "type": "VHS_VideoCombine", "pos": { - "0": 1534, - "1": 490 + "0": 1541, + "1": 339 }, "size": [ 1698.6201171875, 1331.1720703125 ], "flags": {}, - "order": 4, + "order": 6, "mode": 0, "inputs": [ { "name": "images", "type": "IMAGE", - "link": 38 + "link": 53 }, { "name": "audio", @@ -257,7 +174,7 @@ "hidden": false, "paused": false, "params": { - "filename": "PyramidFlow_00061.mp4", + "filename": "PyramidFlow_00089.mp4", "subfolder": "", "type": "output", "format": "video/h264-mp4", @@ -266,6 +183,140 @@ "muted": false } } + }, + { + "id": 36, + "type": "PyramidFlowTextEncodeComfy", + "pos": { + "0": 597, + "1": 779 + }, + "size": { + "0": 400, + "1": 200 + }, + "flags": {}, + "order": 3, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 47 + } + ], + "outputs": [ + { + "name": "prompt_embeds", + "type": "PYRAMIDFLOWPROMPT", + "links": [ + 54 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "PyramidFlowTextEncodeComfy" + }, + "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": 37, + "type": "DualCLIPLoader", + "pos": { + "0": 132, + "1": 780 + }, + "size": [ + 407.1675593807479, + 106 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "CLIP", + "type": "CLIP", + "links": [ + 47 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DualCLIPLoader" + }, + "widgets_values": [ + "clip_l.safetensors", + "t5\\t5xxl_fp16.safetensors", + "flux" + ] + }, + { + "id": 39, + "type": "Note", + "pos": { + "0": 204, + "1": 946 + }, + "size": [ + 318.2556676190985, + 66.48251931043842 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "fp8 text encoder results are different from fp16!" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 5, + "type": "DownloadAndLoadPyramidFlowModel", + "pos": { + "0": 143, + "1": 489 + }, + "size": { + "0": 385.7839050292969, + "1": 202 + }, + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "pyramidflow_model", + "type": "PYRAMIDFLOWMODEL", + "links": [ + 7 + ], + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadPyramidFlowModel" + }, + "widgets_values": [ + "rain1011/pyramid-flow-miniflux", + "diffusion_transformer_384p", + "bf16", + "bf16", + "bf16", + false + ] } ], "links": [ @@ -294,38 +345,38 @@ "LATENT" ], [ - 30, - 5, + 47, + 37, 0, - 22, + 36, 0, - "PYRAMIDFLOWMODEL" + "CLIP" ], [ - 31, - 22, - 0, - 9, - 1, - "PYRAMIDFLOWPROMPT" - ], - [ - 38, + 53, 8, 0, 14, 0, "IMAGE" + ], + [ + 54, + 36, + 0, + 9, + 1, + "PYRAMIDFLOWPROMPT" ] ], "groups": [], "config": {}, "extra": { "ds": { - "scale": 0.6934334949442617, + "scale": 0.6303940863129696, "offset": [ - -267.34972182737584, - -351.34162515690946 + 274.1852517840429, + -178.2662230728557 ] } }, diff --git a/nodes.py b/nodes.py index 654ff94..681bca5 100644 --- a/nodes.py +++ b/nodes.py @@ -64,9 +64,12 @@ class DownloadAndLoadPyramidFlowModel: variant_path = os.path.join(model_path, variant) if not os.path.exists(variant_path): + from huggingface_hub import snapshot_download log.info(f"Downloading model to: {model_path}") + ignore_patterns = [] + if model == "rain1011/pyramid-flow-miniflux": + ignore_patterns.append["*text_encoder*", "*tokenizer*"] if variant == "diffusion_transformer_384p": - from huggingface_hub import snapshot_download snapshot_download( repo_id=model, ignore_patterns=["*diffusion_transformer_768p*"], @@ -74,7 +77,6 @@ class DownloadAndLoadPyramidFlowModel: local_dir_use_symlinks=False, ) elif variant == "diffusion_transformer_768p": - from huggingface_hub import snapshot_download snapshot_download( repo_id=model, ignore_patterns=["*diffusion_transformer_384p*"], @@ -298,6 +300,7 @@ class PyramidFlowTextEncodeComfy: "clip": ("CLIP",), "positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ), "negative_prompt": ("STRING", {"default": "", "multiline": True} ), + "force_offload": ("BOOLEAN", {"default": True}), } } @@ -306,42 +309,36 @@ class PyramidFlowTextEncodeComfy: FUNCTION = "process" CATEGORY = "CogVideoWrapper" - def process(self, clip, positive_prompt, negative_prompt): + def process(self, clip, positive_prompt, negative_prompt, force_offload=True): + max_lenght = 128 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.tokenizer.t5xxl.max_length = max_lenght + clip.tokenizer.t5xxl.min_length = 1 + clip.tokenizer.clip_l.max_length = 77 clip.cond_stage_model.t5xxl.return_attention_masks = True + clip.cond_stage_model.t5xxl.enable_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) + clip.cond_stage_model.to(device)#.to(torch.bfloat16) + clip.cond_stage_model.clip_l.to(device) - max_length = prompt_attention_mask["attention_mask"].shape[1] - prompt_embeds = prompt_embeds[:, :max_length, :] + #positive + tokens = clip.tokenizer.t5xxl.tokenize_with_weights(positive_prompt, return_word_ids=False) + prompt_embeds, _, prompt_attention_mask = clip.cond_stage_model.t5xxl.encode_token_weights(tokens) + tokens = clip.tokenizer.clip_l.tokenize_with_weights(positive_prompt, return_word_ids=False) + _, pooled_prompt_embeds, = clip.cond_stage_model.clip_l.encode_token_weights(tokens) + #negative + tokens = clip.tokenizer.t5xxl.tokenize_with_weights(negative_prompt, return_word_ids=False) + negative_prompt_embeds, _, negative_prompt_attention_mask = clip.cond_stage_model.t5xxl.encode_token_weights(tokens) + tokens = clip.tokenizer.clip_l.tokenize_with_weights(negative_prompt, return_word_ids=False) + _, pooled_negative_prompt_embeds, = clip.cond_stage_model.clip_l.encode_token_weights(tokens) - 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) + if force_offload: + clip.cond_stage_model.to(offload_device) embeds = { "prompt_embeds": prompt_embeds.to(device), @@ -349,7 +346,7 @@ class PyramidFlowTextEncodeComfy: "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) + "negative_pooled_embeds": pooled_negative_prompt_embeds.to(device), } return (embeds, ) @@ -458,7 +455,7 @@ NODE_CLASS_MAPPINGS = { "PyramidFlowVAEDecode": PyramidFlowVAEDecode, "PyramidFlowTextEncode": PyramidFlowTextEncode, "PyramidFlowVAEEncode": PyramidFlowVAEEncode, - #"PyramidFlowTextEncodeComfy": PyramidFlowTextEncodeComfy, + "PyramidFlowTextEncodeComfy": PyramidFlowTextEncodeComfy, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -467,5 +464,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "PyramidFlowVAEDecode" : "PyramidFlow VAE Decode", "PyramidFlowTextEncode": "PyramidFlow Text Encode", "PyramidFlowVAEEncode": "PyramidFlow VAE Encode", - #"PyramidFlowTextEncodeComfy": "PyramidFlow Text Encode Comfy", + "PyramidFlowTextEncodeComfy": "PyramidFlow Text Encode Comfy", } diff --git a/pyramid_dit/flux_modules/modeling_text_encoder.py b/pyramid_dit/flux_modules/modeling_text_encoder.py index aaffbb2..82bf56d 100644 --- a/pyramid_dit/flux_modules/modeling_text_encoder.py +++ b/pyramid_dit/flux_modules/modeling_text_encoder.py @@ -124,6 +124,13 @@ class FluxTextEncoderWithMask(nn.Module): num_images_per_prompt=num_images_per_prompt, device=device, ) + print("prompt_embeds_shape: ",prompt_embeds.shape) + print("pooled_prompt_embeds_shape: ",pooled_prompt_embeds.shape) + print("prompt_attention_mask_shape: ",prompt_attention_mask.shape) + # prompt_embeds_shape: torch.Size([1, 128, 4096]) + # pooled_prompt_embeds_shape: torch.Size([1, 768]) + # prompt_attention_mask_shape: torch.Size([1, 128]) + return prompt_embeds, prompt_attention_mask, pooled_prompt_embeds diff --git a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py index d77f7ee..b9bf1f4 100644 --- a/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py +++ b/pyramid_dit/pyramid_dit_for_video_gen_pipeline.py @@ -131,11 +131,10 @@ class PyramidDiTForVideoGeneration: use_temporal_causal=use_temporal_causal, ) - if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: for name, param in self.dit.named_parameters(): - if name != "pos_embedding": - param.data = param.data.to(model_dtype) + if name != "pos_embedding": + param.data = param.data.to(model_dtype) if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2] and fp8_fastmode: from ..fp8_optimization import convert_fp8_linear