Big update, refactor model loading, breaks backwards compatibility

archiving old  version in legacy branch, probably not going to support SD3 version going forward
This commit is contained in:
kijai
2024-10-30 15:48:32 +02:00
parent 00e2dcb632
commit 3e14b0d77e
10 changed files with 1125 additions and 1869 deletions
+92
View File
@@ -0,0 +1,92 @@
{
"_class_name": "CausalVideoVAE",
"_diffusers_version": "0.29.2",
"add_post_quant_conv": true,
"decoder_act_fn": "silu",
"decoder_block_dropout": [
0.0,
0.0,
0.0,
0.0
],
"decoder_block_out_channels": [
128,
256,
512,
512
],
"decoder_in_channels": 16,
"decoder_layers_per_block": [
3,
3,
3,
3
],
"decoder_norm_num_groups": 32,
"decoder_out_channels": 3,
"decoder_spatial_up_sample": [
true,
true,
true,
false
],
"decoder_temporal_up_sample": [
true,
true,
true,
false
],
"decoder_type": "causal_vae_conv",
"decoder_up_block_types": [
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D",
"UpDecoderBlockCausal3D"
],
"downsample_scale": 8,
"encoder_act_fn": "silu",
"encoder_block_dropout": [
0.0,
0.0,
0.0,
0.0
],
"encoder_block_out_channels": [
128,
256,
512,
512
],
"encoder_double_z": true,
"encoder_down_block_types": [
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D",
"DownEncoderBlockCausal3D"
],
"encoder_in_channels": 3,
"encoder_layers_per_block": [
2,
2,
2,
2
],
"encoder_norm_num_groups": 32,
"encoder_out_channels": 16,
"encoder_spatial_down_sample": [
true,
true,
true,
false
],
"encoder_temporal_down_sample": [
true,
true,
true,
false
],
"encoder_type": "causal_vae_conv",
"interpolate": false,
"sample_size": 256,
"scaling_factor": 0.13025
}
+21
View File
@@ -0,0 +1,21 @@
{
"_class_name": "PyramidFluxTransformer",
"_diffusers_version": "0.30.3",
"attention_head_dim": 64,
"axes_dims_rope": [
16,
24,
24
],
"in_channels": 64,
"interp_condition_pos": true,
"joint_attention_dim": 4096,
"num_attention_heads": 30,
"num_layers": 8,
"num_single_layers": 16,
"patch_size": 1,
"pooled_projection_dim": 768,
"use_flash_attn": false,
"use_gradient_checkpointing": false,
"use_temporal_causal": true
}
+20
View File
@@ -0,0 +1,20 @@
{
"_class_name": "PyramidDiffusionMMDiT",
"_diffusers_version": "0.30.0",
"attention_head_dim": 64,
"caption_projection_dim": 1536,
"in_channels": 16,
"joint_attention_dim": 4096,
"max_num_frames": 200,
"num_attention_heads": 24,
"num_layers": 24,
"patch_size": 2,
"pooled_projection_dim": 2048,
"pos_embed_max_size": 192,
"pos_embed_type": "sincos",
"qk_norm": "rms_norm",
"sample_size": 128,
"use_flash_attn": false,
"use_gradient_checkpointing": false,
"use_temporal_causal": true
}
@@ -1,31 +1,133 @@
{ {
"last_node_id": 39, "last_node_id": 53,
"last_link_id": 54, "last_link_id": 81,
"nodes": [ "nodes": [
{ {
"id": 9, "id": 39,
"type": "PyramidFlowSampler", "type": "Note",
"pos": { "pos": {
"0": 1059, "0": 30,
"1": 497 "1": 650
}, },
"size": { "size": {
"0": 411.5168151855469, "0": 318.25567626953125,
"1": 66.4825210571289
},
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"fp8 text encoder results are different from fp16!"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 48,
"type": "PyramidFlowVAEDecode",
"pos": {
"0": 1051,
"1": 534
},
"size": {
"0": 315,
"1": 126
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "vae",
"type": "PYRAMIDFLOWVAE",
"link": 71
},
{
"name": "samples",
"type": "LATENT",
"link": 76
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"links": [
77
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowVAEDecode"
},
"widgets_values": [
256,
2,
true
]
},
{
"id": 37,
"type": "DualCLIPLoader",
"pos": {
"0": -40,
"1": 480
},
"size": {
"0": 407.1675720214844,
"1": 106
},
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "CLIP",
"type": "CLIP",
"links": [
80
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "DualCLIPLoader"
},
"widgets_values": [
"clip_l.safetensors",
"t5\\t5xxl_fp16.safetensors",
"flux"
]
},
{
"id": 50,
"type": "PyramidFlowSampler",
"pos": {
"0": 1046,
"1": 144
},
"size": {
"0": 315,
"1": 314 "1": 314
}, },
"flags": {}, "flags": {},
"order": 4, "order": 5,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
"name": "model", "name": "model",
"type": "PYRAMIDFLOWMODEL", "type": "PYRAMIDFLOWMODEL",
"link": 7 "link": 74
}, },
{ {
"name": "prompt_embeds", "name": "prompt_embeds",
"type": "PYRAMIDFLOWPROMPT", "type": "PYRAMIDFLOWPROMPT",
"link": 54 "link": 81
}, },
{ {
"name": "input_latent", "name": "input_latent",
@@ -35,20 +137,12 @@
} }
], ],
"outputs": [ "outputs": [
{
"name": "model",
"type": "PYRAMIDFLOWMODEL",
"links": [
8
]
},
{ {
"name": "samples", "name": "samples",
"type": "LATENT", "type": "LATENT",
"links": [ "links": [
9 76
], ]
"slot_index": 1
} }
], ],
"properties": { "properties": {
@@ -60,76 +154,152 @@
"20, 20, 20", "20, 20, 20",
"10, 10, 10", "10, 10, 10",
16, 16,
9, 7,
5, 5,
44664248661395, 44664248661402,
"fixed", "fixed",
"" false
] ]
}, },
{ {
"id": 8, "id": 40,
"type": "PyramidFlowVAEDecode", "type": "PyramidFlowTransformerLoader",
"pos": { "pos": {
"0": 1161, "0": 225,
"1": 873 "1": 140
}, },
"size": { "size": {
"0": 315, "0": 444.05462646484375,
"1": 102 "1": 82
}, },
"flags": {}, "flags": {},
"order": 5, "order": 2,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
"name": "model", "name": "compile_args",
"type": "PYRAMIDFLOWMODEL", "type": "MOCHICOMPILEARGS",
"link": 8 "link": null,
}, "shape": 7
{
"name": "samples",
"type": "LATENT",
"link": 9
} }
], ],
"outputs": [ "outputs": [
{ {
"name": "images", "name": "pyramidflow_model",
"type": "IMAGE", "type": "PYRAMIDFLOWMODEL",
"links": [ "links": [
53 74
], ],
"slot_index": 0 "slot_index": 0
} }
], ],
"properties": { "properties": {
"Node name for S&R": "PyramidFlowVAEDecode" "Node name for S&R": "PyramidFlowTransformerLoader"
}, },
"widgets_values": [ "widgets_values": [
256, "pyramidflow\\pyramid_flow_miniflux_bf16_v2.safetensors",
2 "bf16"
] ]
}, },
{ {
"id": 14, "id": 43,
"type": "VHS_VideoCombine", "type": "PyramidFlowVAELoader",
"pos": { "pos": {
"0": 1541, "0": 250,
"1": 339 "1": 282
},
"size": {
"0": 411.12652587890625,
"1": 82
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"type": "MOCHICOMPILEARGS",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "pyramidflow_vae",
"type": "PYRAMIDFLOWVAE",
"links": [
71
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowVAELoader"
},
"widgets_values": [
"pyramidflow\\pyramid_flow_vae_bf16.safetensors",
"bf16"
]
},
{
"id": 53,
"type": "PyramidFlowTextEncode",
"pos": {
"0": 444,
"1": 476
}, },
"size": [ "size": [
1698.6201171875, 437.19819084177436,
1331.1720703125 269.9795836111515
], ],
"flags": {}, "flags": {},
"order": 6, "order": 4,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 80
}
],
"outputs": [
{
"name": "prompt_embeds",
"type": "PYRAMIDFLOWPROMPT",
"links": [
81
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowTextEncode"
},
"widgets_values": [
"Beautiful, snowy Tokyo city is bustling. The camera moves through the bustling city street, following several people enjoying the beautiful snowy weather and shopping at nearby stalls. Gorgeous sakura petals are flying through the wind along with snowflakes, 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": 51,
"type": "VHS_VideoCombine",
"pos": {
"0": 1420,
"1": 113
},
"size": [
1018.2306485015952,
922.9384155273438
],
"flags": {},
"order": 7,
"mode": 0, "mode": 0,
"inputs": [ "inputs": [
{ {
"name": "images", "name": "images",
"type": "IMAGE", "type": "IMAGE",
"link": 53 "link": 77
}, },
{ {
"name": "audio", "name": "audio",
@@ -174,7 +344,7 @@
"hidden": false, "hidden": false,
"paused": false, "paused": false,
"params": { "params": {
"filename": "PyramidFlow_00089.mp4", "filename": "PyramidFlow_00129.mp4",
"subfolder": "", "subfolder": "",
"type": "output", "type": "output",
"format": "video/h264-mp4", "format": "video/h264-mp4",
@@ -183,188 +353,54 @@
"muted": false "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": [ "links": [
[ [
7, 71,
5, 43,
0, 0,
9, 48,
0,
"PYRAMIDFLOWVAE"
],
[
74,
40,
0,
50,
0, 0,
"PYRAMIDFLOWMODEL" "PYRAMIDFLOWMODEL"
], ],
[ [
8, 76,
9, 50,
0, 0,
8, 48,
0,
"PYRAMIDFLOWMODEL"
],
[
9,
9,
1,
8,
1, 1,
"LATENT" "LATENT"
], ],
[ [
47, 77,
37, 48,
0, 0,
36, 51,
0,
"CLIP"
],
[
53,
8,
0,
14,
0, 0,
"IMAGE" "IMAGE"
], ],
[ [
54, 80,
36, 37,
0, 0,
9, 53,
0,
"CLIP"
],
[
81,
53,
0,
50,
1, 1,
"PYRAMIDFLOWPROMPT" "PYRAMIDFLOWPROMPT"
] ]
@@ -373,10 +409,10 @@
"config": {}, "config": {},
"extra": { "extra": {
"ds": { "ds": {
"scale": 0.6303940863129696, "scale": 0.7627768444386932,
"offset": [ "offset": [
274.1852517840429, 205.40187889258976,
-178.2662230728557 148.78644149639825
] ]
} }
}, },
@@ -1,333 +0,0 @@
{
"last_node_id": 27,
"last_link_id": 38,
"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": 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-sd3",
"diffusion_transformer_768p",
"bf16",
"bf16",
"bf16",
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": [
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": 9,
"type": "PyramidFlowSampler",
"pos": {
"0": 1059,
"1": 497
},
"size": {
"0": 411.5168151855469,
"1": 314
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "PYRAMIDFLOWMODEL",
"link": 7
},
{
"name": "prompt_embeds",
"type": "PYRAMIDFLOWPROMPT",
"link": 31
},
{
"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,
9,
5,
44664248661394,
"fixed",
""
]
},
{
"id": 14,
"type": "VHS_VideoCombine",
"pos": {
"0": 1534,
"1": 490
},
"size": [
1698.6201171875,
1331.1720703125
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 38
},
{
"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": 24,
"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_00060.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
},
"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"
],
[
31,
22,
0,
9,
1,
"PYRAMIDFLOWPROMPT"
],
[
38,
8,
0,
14,
0,
"IMAGE"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.6830134553650706,
"offset": [
-385.88395293212943,
-310.74071933624276
]
}
},
"version": 0.4
}
@@ -1,457 +0,0 @@
{
"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
}
+227 -208
View File
@@ -1,6 +1,7 @@
import os import os
import torch import torch
import folder_paths import folder_paths
import json
import comfy.model_management as mm import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file from comfy.utils import ProgressBar, load_torch_file
@@ -17,29 +18,115 @@ script_directory = os.path.dirname(os.path.abspath(__file__))
if not "pyramidflow" in folder_paths.folder_names_and_paths: if not "pyramidflow" in folder_paths.folder_names_and_paths:
folder_paths.add_model_folder_path("pyramidflow", os.path.join(folder_paths.models_dir, "pyramidflow")) folder_paths.add_model_folder_path("pyramidflow", os.path.join(folder_paths.models_dir, "pyramidflow"))
class DownloadAndLoadPyramidFlowModel: from .pyramid_dit.mmdit_modules import PyramidDiffusionMMDiT
from .pyramid_dit.flux_modules import PyramidFluxTransformer
from .video_vae.modeling_causal_vae import CausalVideoVAE
from contextlib import nullcontext
try:
from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
is_accelerate_available = True
except:
is_accelerate_available = False
class PyramidFlowTorchCompileSettings:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"model": ( "backend": (["inductor","cudagraphs"], {"default": "inductor"}),
[ "fullgraph": ("BOOLEAN", {"default": False, "tooltip": "Enable full graph mode"}),
"rain1011/pyramid-flow-sd3", "mode": (["default", "max-autotune", "max-autotune-no-cudagraphs", "reduce-overhead"], {"default": "default"}),
"rain1011/pyramid-flow-miniflux" "compile_whole_model": ("BOOLEAN", {"default": False, "tooltip": "Compile the whole model, overrides other block settings"}),
"single_blocks": ("BOOLEAN", {"default": True, "tooltip": "Compile single_blocks"}),
"double_blocks": ("BOOLEAN", {"default": True, "tooltip": "Compile transformer blocks"}),
"embedders": ("BOOLEAN", {"default": True, "tooltip": "Compile embedders"}),
"compile_rest": ("BOOLEAN", {"default": True, "tooltip": "Compile the rest of the model (proj and norm out)"}),
},
}
RETURN_TYPES = ("MOCHICOMPILEARGS",)
RETURN_NAMES = ("torch_compile_args",)
FUNCTION = "loadmodel"
CATEGORY = "MochiWrapper"
DESCRIPTION = "torch.compile settings, when connected to the model loader, torch.compile of the selected layers is attempted. Requires Triton and torch 2.5.0 is recommended"
], def loadmodel(self, backend, fullgraph, mode, compile_whole_model, single_blocks, double_blocks, embedders, compile_rest):
),
"variant": (
["diffusion_transformer_384p", "diffusion_transformer_768p"],
),
compile_args = {
"backend": backend,
"fullgraph": fullgraph,
"mode": mode,
"compile_whole_model": compile_whole_model,
"single_blocks": single_blocks,
"double_blocks": double_blocks,
"embedders": embedders,
"compile_rest": compile_rest,
}
return (compile_args, )
class PyramidFlowVAELoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"vae": (folder_paths.get_filename_list("vae"), {"tooltip": "The name of the checkpoint (model) to load.",}),
"precision": (["fp16", "bf16", "fp32"], {"default": "bf16"}),
}, },
"optional": { "optional": {
"model_dtype": (["fp8_e4m3fn","fp8_e5m2","fp16", "fp32", "bf16"],{"default": "bf16", }), "compile_args": ("MOCHICOMPILEARGS", {"tooltip": "Optional torch.compile arguments",}),
"text_encoder_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), }
"vae_dtype": (["fp16", "fp32", "bf16"],{"default": "bf16", }), }
"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"}), RETURN_TYPES = ("PYRAMIDFLOWVAE", )
RETURN_NAMES = ("pyramidflow_vae",)
FUNCTION = "loadmodel"
CATEGORY = "PyramidFlowWrapper"
def loadmodel(self, vae, precision, compile_args=None):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
vae_path = folder_paths.get_full_path_or_raise("vae", vae)
dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
config_path = os.path.join(script_directory, 'configs', 'causal_video_vae_config.json')
with open(config_path) as f:
config = json.load(f)
with (init_empty_weights() if is_accelerate_available else nullcontext()):
vae = CausalVideoVAE.from_config(config, torch_dtype=dtype, interpolate=False)
vae_sd = load_torch_file(vae_path)
if is_accelerate_available:
for name, param in vae.named_parameters():
set_module_tensor_to_device(vae, name, dtype=dtype, device=device, value=vae_sd[name])
else:
vae.load_state_dict(vae_sd)
del vae_sd
# Freeze vae
for parameter in vae.parameters():
parameter.requires_grad = False
vae.eval().to(device)
#torch.compile
if compile_args is not None:
vae = torch.compile(vae, fullgraph=compile_args["fullgraph"], dynamic=False, backend=compile_args["backend"])
return (vae,)
class PyramidFlowModelLoader:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "The name of the checkpoint (model) to load.",}),
"precision": (["fp8_e4m3fn","fp8_e4m3fn_fast","fp16", "fp32", "bf16"], {"default": "bf16"}),
},
"optional": {
"compile_args": ("MOCHICOMPILEARGS", {"tooltip": "Optional torch.compile arguments",}),
} }
} }
@@ -48,111 +135,90 @@ class DownloadAndLoadPyramidFlowModel:
FUNCTION = "loadmodel" FUNCTION = "loadmodel"
CATEGORY = "PyramidFlowWrapper" CATEGORY = "PyramidFlowWrapper"
def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, fp8_fastmode): def loadmodel(self, model, precision, compile_args=None):
device = mm.get_torch_device() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
mm.soft_empty_cache()
model_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32, "fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e5m2": torch.float8_e5m2}[model_dtype] model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
text_encoder_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[text_encoder_dtype]
vae_dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[vae_dtype]
base_path = folder_paths.get_folder_paths("pyramidflow")[0] dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
model_path = os.path.join(base_path, model.split("/")[-1]) transformer_sd = load_torch_file(model_path)
variant_path = os.path.join(model_path, variant)
if not os.path.exists(variant_path): for key in transformer_sd:
from huggingface_hub import snapshot_download if key.startswith("pos_embed."):
log.info(f"Downloading model to: {model_path}") model_name = "pyramid_mmdit"
ignore_patterns = [] continue
if model == "rain1011/pyramid-flow-miniflux": else:
ignore_patterns.extend["*text_encoder*", "*tokenizer*"] model_name = "pyramid_flux"
if variant == "diffusion_transformer_384p":
snapshot_download(
repo_id=model,
ignore_patterns=["*diffusion_transformer_768p*"],
local_dir=model_path,
local_dir_use_symlinks=False,
)
elif variant == "diffusion_transformer_768p":
snapshot_download(
repo_id=model,
ignore_patterns=["*diffusion_transformer_384p*"],
local_dir=model_path,
local_dir_use_symlinks=False,
)
model_name = "pyramid_flux" if "flux" in model else "pyramid_mmdit"
print(model_name)
model = PyramidDiTForVideoGeneration(
model_path,
model_dtype,
model_name,
text_encoder_dtype,
vae_dtype,
model_variant=variant,
fp8_fastmode=fp8_fastmode,
)
# # compilation if model_name == "pyramid_flux":
# if compile == "torch": config_path = os.path.join(script_directory, 'configs', 'miniflux_transformer_config.json')
# torch._dynamo.config.suppress_errors = True with open(config_path) as f:
# pipe.transformer.to(memory_format=torch.channels_last) config = json.load(f)
# pipe.transformer = torch.compile(pipe.transformer, mode="max-autotune", fullgraph=True)
# elif compile == "onediff":
# from onediffx import compile_pipe
# os.environ['NEXFORT_FX_FORCE_TRITON_SDPA'] = '1'
# pipe = compile_pipe( with (init_empty_weights() if is_accelerate_available else nullcontext()):
# pipe, transformer = PyramidFluxTransformer.from_config(config)
# backend="nexfort",
# options= {"mode": "max-optimize:max-autotune:max-autotune", "memory_format": "channels_last", "options": {"inductor.optimize_linear_epilogue": False, "triton.fuse_attention_allow_fp16_reduction": False}},
# ignores=["vae"],
# fuse_qkv_projections=True if pab_config is None else False,
# )
pyramid_pipe = { if is_accelerate_available:
"model": model, logging.info("Using accelerate to load and assign model weights to device...")
"dtype": model_dtype, for name, param in transformer.named_parameters():
"text_encoder_dtype": text_encoder_dtype, set_module_tensor_to_device(transformer, name, dtype=dtype, device=device, value=transformer_sd[name])
"vae_dtype": vae_dtype, else:
} transformer.load_state_dict(transformer_sd)
return (pyramid_pipe,) transformer = transformer.to(dtype)
elif model_name == "pyramid_mmdit":
config_path = os.path.join(script_directory, 'configs', 'mmdit_transformer_config.json')
with open(config_path) as f:
config = json.load(f)
transformer = PyramidDiffusionMMDiT.from_config(config)
params_to_keep = {"pos_embedding"}
if is_accelerate_available:
logging.info("Using accelerate to load and assign model weights to device...")
for name, param in transformer.named_parameters():
if not any(keyword in name for keyword in params_to_keep):
set_module_tensor_to_device(transformer, name, dtype=dtype, device=device, value=transformer_sd[name])
else:
set_module_tensor_to_device(transformer, name, dtype=torch.bfloat16, device=device, value=transformer_sd[name])
else:
transformer.load_state_dict(transformer_sd)
if dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
for name, param in transformer.named_parameters():
if not any(keyword in name for keyword in params_to_keep):
param.data = param.data.to(dtype)
if precision == "fp8_e4m3fn_fast":
from .fp8_optimization import convert_fp8_linear
convert_fp8_linear(transformer, torch.bfloat16)
transformer.to(device)
#torch.compile
if compile_args is not None:
torch._dynamo.config.force_parameter_static_shapes = False
dynamic = True # because of the stages the compiliation should be dynamic
if compile_args["compile_whole_model"]:
transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
else:
if compile_args["single_blocks"]:
for i, block in enumerate(transformer.single_transformer_blocks):
transformer.single_transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
if compile_args["double_blocks"]:
for i, block in enumerate(transformer.transformer_blocks):
transformer.transformer_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
if compile_args["embedders"]:
transformer.context_embedder = torch.compile(transformer.context_embedder, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
transformer.time_text_embed = torch.compile(transformer.time_text_embed, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
transformer.x_embedder = torch.compile(transformer.x_embedder, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
if compile_args["compile_rest"]:
transformer.norm_out.linear = torch.compile(transformer.norm_out.linear, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
transformer.proj_out = torch.compile(transformer.proj_out, fullgraph=compile_args["fullgraph"], dynamic=dynamic, backend=compile_args["backend"])
class CogVideoTextEncode: pyramid_model = PyramidDiTForVideoGeneration(transformer, dtype, model_name)
@classmethod
def INPUT_TYPES(s):
return {"required": {
"clip": ("CLIP",),
"prompt": ("STRING", {"default": "", "multiline": True} ),
},
"optional": {
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"force_offload": ("BOOLEAN", {"default": True}),
}
}
RETURN_TYPES = ("CONDITIONING",) return (pyramid_model,)
RETURN_NAMES = ("conditioning",)
FUNCTION = "process"
CATEGORY = "CogVideoWrapper"
def process(self, clip, prompt, strength=1.0, force_offload=True):
load_device = mm.text_encoder_device()
offload_device = mm.text_encoder_offload_device()
clip.tokenizer.t5xxl.pad_to_max_length = True
clip.tokenizer.t5xxl.max_length = 226
clip.cond_stage_model.to(load_device)
tokens = clip.tokenize(prompt, return_word_ids=True)
embeds = clip.encode_from_tokens(tokens, return_pooled=False, return_dict=False)
embeds *= strength
if force_offload:
clip.cond_stage_model.to(offload_device)
return (embeds, )
class PyramidFlowSampler: class PyramidFlowSampler:
@classmethod @classmethod
@@ -177,8 +243,8 @@ class PyramidFlowSampler:
} }
} }
RETURN_TYPES = ("PYRAMIDFLOWMODEL", "LATENT", ) RETURN_TYPES = ("LATENT", )
RETURN_NAMES = ("model","samples", ) RETURN_NAMES = ("samples", )
FUNCTION = "sample" FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper" CATEGORY = "PyramidFlowWrapper"
@@ -188,7 +254,14 @@ class PyramidFlowSampler:
device = mm.get_torch_device() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
dtype = model["dtype"]
if isinstance(model, dict):
pyramid_model = model["model"]
#dtype = model["dtype"]
else:
pyramid_model = model
dtype = pyramid_model.dit.dtype
torch.manual_seed(seed) torch.manual_seed(seed)
torch.cuda.manual_seed(seed) torch.cuda.manual_seed(seed)
@@ -202,7 +275,7 @@ class PyramidFlowSampler:
if input_latent is None: if input_latent is None:
with autocast_context: with autocast_context:
latents = model["model"].generate( latents = pyramid_model.generate(
prompt_embeds_dict = prompt_embeds, prompt_embeds_dict = prompt_embeds,
device=device, device=device,
num_inference_steps=first_frame_steps, num_inference_steps=first_frame_steps,
@@ -216,11 +289,11 @@ class PyramidFlowSampler:
) )
else: else:
with autocast_context: with autocast_context:
latents = model["model"].generate_i2v( latents = pyramid_model.generate_i2v(
prompt_embeds_dict = prompt_embeds, prompt_embeds_dict = prompt_embeds,
input_image_latent=input_latent, input_image_latent=input_latent,
device=device, device=device,
num_inference_steps=video_steps, #why's this a list num_inference_steps=video_steps,
height=height, height=height,
width=width, width=width,
temp=temp, temp=temp,
@@ -228,78 +301,19 @@ class PyramidFlowSampler:
output_type="latent", output_type="latent",
) )
if not keep_model_loaded: if not keep_model_loaded:
model["model"].dit.to(offload_device) pyramid_model.dit.to(offload_device)
return (model, {"samples": latents},) return ({"samples": latents},)
#todo: sd3 version
class PyramidFlowTextEncode: class PyramidFlowTextEncode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("PYRAMIDFLOWMODEL",),
"positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ),
"keep_model_loaded": ("BOOLEAN", {"default": False}),
},
"optional": {
"prev_prompt": ("PYRAMIDFLOWPROMPT", ),
}
}
RETURN_TYPES = ("PYRAMIDFLOWPROMPT", )
RETURN_NAMES = ("prompt_embeds", )
FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper"
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()
text_encoder = model["model"].text_encoder
autocastcondition = not model["text_encoder_dtype"] == torch.float32
autocast_context = torch.autocast(mm.get_autocast_device(device), dtype=model["text_encoder_dtype"]) if autocastcondition else nullcontext()
text_encoder.to(device)
with autocast_context:
prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = text_encoder(positive_prompt, device)
negative_prompt_embeds, negative_prompt_attention_mask, pooled_negative_prompt_embeds = text_encoder(negative_prompt, device)
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,
"pooled_embeds": pooled_prompt_embeds,
"negative_prompt_embeds": negative_prompt_embeds,
"negative_attention_mask": negative_prompt_attention_mask,
"negative_pooled_embeds": pooled_negative_prompt_embeds
}
return (embeds,)
#not functional yet, todo: figure out why the results are bad with it
class PyramidFlowTextEncodeComfy:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { return {"required": {
"clip": ("CLIP",), "clip": ("CLIP",),
"positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ), "positive_prompt": ("STRING", {"default": "hyper quality, Ultra HD, 8K", "multiline": True} ),
"negative_prompt": ("STRING", {"default": "", "multiline": True} ), "negative_prompt": ("STRING", {"default": "cartoon style, worst quality, low quality, blurry, absolute black, absolute white, low res, extra limbs, extra digits, misplaced objects, mutated anatomy, monochrome, horror", "multiline": True} ),
"force_offload": ("BOOLEAN", {"default": True}), "force_offload": ("BOOLEAN", {"default": True}),
} }
} }
@@ -324,7 +338,6 @@ class PyramidFlowTextEncodeComfy:
clip.cond_stage_model.t5_attention_mask = True clip.cond_stage_model.t5_attention_mask = True
clip.cond_stage_model.to(device)#.to(torch.bfloat16) clip.cond_stage_model.to(device)#.to(torch.bfloat16)
clip.cond_stage_model.clip_l.to(device)
#positive #positive
tokens = clip.tokenizer.t5xxl.tokenize_with_weights(positive_prompt, return_word_ids=False) tokens = clip.tokenizer.t5xxl.tokenize_with_weights(positive_prompt, return_word_ids=False)
@@ -355,8 +368,9 @@ class PyramidFlowVAEEncode:
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"model": ("PYRAMIDFLOWMODEL",), "vae": ("PYRAMIDFLOWVAE",),
"image": ("IMAGE",), "image": ("IMAGE",),
"enable_tiling": ("BOOLEAN", {"default": False}),
}, },
} }
@@ -365,19 +379,20 @@ class PyramidFlowVAEEncode:
FUNCTION = "sample" FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper" CATEGORY = "PyramidFlowWrapper"
def sample(self, model, image): def sample(self, vae, image, enable_tiling):
mm.soft_empty_cache() mm.soft_empty_cache()
self.vae = model["model"].vae dtype = vae.dtype
dtype = model["vae_dtype"] if enable_tiling:
vae.enable_tiling()
else:
vae.disable_tiling()
device = mm.get_torch_device() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
self.vae.disable_tiling()
# For the image latent # For the image latent
self.vae_shift_factor = 0.1490 vae_shift_factor = 0.1490
self.vae_scale_factor = 1 / 1.8415 vae_scale_factor = 1 / 1.8415
normalize = transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)) normalize = transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))
input_image_tensor = rearrange(image, 'b h w c -> b c h w') input_image_tensor = rearrange(image, 'b h w c -> b c h w')
@@ -385,10 +400,9 @@ class PyramidFlowVAEEncode:
input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1 input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1
input_image_tensor = input_image_tensor.to(dtype=dtype, device=device) input_image_tensor = input_image_tensor.to(dtype=dtype, device=device)
self.vae.to(device) vae.to(device)
input_image_latent = (self.vae.encode(input_image_tensor).latent_dist.sample() - self.vae_shift_factor) * self.vae_scale_factor # [b c 1 h w] input_image_latent = (vae.encode(input_image_tensor).latent_dist.sample() - vae_shift_factor) * vae_scale_factor # [b c 1 h w]
self.vae.to(offload_device) vae.to(offload_device)
return (input_image_latent,) return (input_image_latent,)
@@ -397,11 +411,11 @@ class PyramidFlowVAEDecode:
def INPUT_TYPES(s): def INPUT_TYPES(s):
return { return {
"required": { "required": {
"model": ("PYRAMIDFLOWMODEL",), "vae": ("PYRAMIDFLOWVAE",),
"samples": ("LATENT",), "samples": ("LATENT",),
"tile_sample_min_size": ("INT", {"default": 256, "min": 64, "max": 512, "step": 8}), "tile_sample_min_size": ("INT", {"default": 256, "min": 64, "max": 512, "step": 8}),
"window_size": ("INT", {"default": 2, "min": 1, "max": 4, "step": 1}), "window_size": ("INT", {"default": 2, "min": 1, "max": 4, "step": 1}),
"enable_tiling": ("BOOLEAN", {"default": True}),
}, },
} }
@@ -410,35 +424,37 @@ class PyramidFlowVAEDecode:
FUNCTION = "sample" FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper" CATEGORY = "PyramidFlowWrapper"
def sample(self, model, samples, tile_sample_min_size, window_size): def sample(self, vae, samples, tile_sample_min_size, window_size, enable_tiling):
mm.soft_empty_cache() mm.soft_empty_cache()
latents = samples["samples"] latents = samples["samples"]
self.vae = model["model"].vae
device = mm.get_torch_device() device = mm.get_torch_device()
offload_device = mm.unet_offload_device() offload_device = mm.unet_offload_device()
self.vae.enable_tiling() if enable_tiling:
vae.enable_tiling()
else:
vae.disable_tiling()
# For the image latent # For the image latent
self.vae_shift_factor = 0.1490 vae_shift_factor = 0.1490
self.vae_scale_factor = 1 / 1.8415 vae_scale_factor = 1 / 1.8415
# For the video latent # For the video latent
self.vae_video_shift_factor = -0.2343 vae_video_shift_factor = -0.2343
self.vae_video_scale_factor = 1 / 3.0986 vae_video_scale_factor = 1 / 3.0986
self.vae.to(device) vae.to(device)
latents = latents.to(self.vae.dtype) latents = latents.to(vae.dtype)
if latents.shape[2] == 1: if latents.shape[2] == 1:
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor latents = (latents / vae_scale_factor) + vae_shift_factor
else: else:
latents[:, :, :1] = (latents[:, :, :1] / self.vae_scale_factor) + self.vae_shift_factor latents[:, :, :1] = (latents[:, :, :1] / vae_scale_factor) + vae_shift_factor
latents[:, :, 1:] = (latents[:, :, 1:] / self.vae_video_scale_factor) + self.vae_video_shift_factor latents[:, :, 1:] = (latents[:, :, 1:] / vae_video_scale_factor) + vae_video_shift_factor
image = self.vae.decode(latents, temporal_chunk=True, window_size=window_size, tile_sample_min_size=tile_sample_min_size).sample image = vae.decode(latents, temporal_chunk=True, window_size=window_size, tile_sample_min_size=tile_sample_min_size).sample
self.vae.to(offload_device) vae.to(offload_device)
image = image.float() image = image.float()
image = (image / 2 + 0.5).clamp(0, 1) image = (image / 2 + 0.5).clamp(0, 1)
@@ -450,12 +466,13 @@ class PyramidFlowVAEDecode:
NODE_CLASS_MAPPINGS = { NODE_CLASS_MAPPINGS = {
"DownloadAndLoadPyramidFlowModel": DownloadAndLoadPyramidFlowModel,
"PyramidFlowSampler": PyramidFlowSampler, "PyramidFlowSampler": PyramidFlowSampler,
"PyramidFlowVAEDecode": PyramidFlowVAEDecode, "PyramidFlowVAEDecode": PyramidFlowVAEDecode,
"PyramidFlowTextEncode": PyramidFlowTextEncode, "PyramidFlowTextEncode": PyramidFlowTextEncode,
"PyramidFlowVAEEncode": PyramidFlowVAEEncode, "PyramidFlowVAEEncode": PyramidFlowVAEEncode,
"PyramidFlowTextEncodeComfy": PyramidFlowTextEncodeComfy, "PyramidFlowTorchCompileSettings": PyramidFlowTorchCompileSettings,
"PyramidFlowTransformerLoader": PyramidFlowModelLoader,
"PyramidFlowVAELoader": PyramidFlowVAELoader
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
@@ -464,5 +481,7 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PyramidFlowVAEDecode" : "PyramidFlow VAE Decode", "PyramidFlowVAEDecode" : "PyramidFlow VAE Decode",
"PyramidFlowTextEncode": "PyramidFlow Text Encode", "PyramidFlowTextEncode": "PyramidFlow Text Encode",
"PyramidFlowVAEEncode": "PyramidFlow VAE Encode", "PyramidFlowVAEEncode": "PyramidFlow VAE Encode",
"PyramidFlowTextEncodeComfy": "PyramidFlow Text Encode Comfy", "PyramidFlowTorchCompileSettings": "PyramidFlow Torch Compile Settings",
"PyramidFlowTransformerLoader": "PyramidFlow Model Loader",
"PyramidFlowVAELoader": "PyramidFlow VAE Loader"
} }
+19 -142
View File
@@ -1,9 +1,6 @@
import torch import torch
import os
import torch.nn.functional as F import torch.nn.functional as F
from collections import OrderedDict
from einops import rearrange from einops import rearrange
from diffusers.utils.torch_utils import randn_tensor from diffusers.utils.torch_utils import randn_tensor
@@ -12,18 +9,7 @@ from tqdm import tqdm
from typing import List, Optional, Union from typing import List, Optional, Union
from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler from ..diffusion_schedulers import PyramidFlowMatchEulerDiscreteScheduler
from ..video_vae.modeling_causal_vae import CausalVideoVAE
from .mmdit_modules import (
PyramidDiffusionMMDiT,
SD3TextEncoderWithMask,
)
from .flux_modules import (
PyramidFluxTransformer,
FluxTextEncoderWithMask,
)
from comfy.utils import ProgressBar from comfy.utils import ProgressBar
def compute_density_for_timestep_sampling( def compute_density_for_timestep_sampling(
@@ -40,74 +26,32 @@ def compute_density_for_timestep_sampling(
u = torch.rand(size=(batch_size,), device="cpu") u = torch.rand(size=(batch_size,), device="cpu")
return u return u
def build_pyramid_dit(
model_name : str,
model_path : str,
torch_dtype,
use_flash_attn : bool,
#use_mixed_training: bool,
interp_condition_pos: bool = True,
use_gradient_checkpointing: bool = False,
use_temporal_causal: bool = True,
gradient_checkpointing_ratio: float = 0.6,
):
#model_dtype = torch.float32 if use_mixed_training else torch_dtype
if model_name == "pyramid_flux":
dit = PyramidFluxTransformer.from_pretrained(
model_path, torch_dtype=torch_dtype,
use_gradient_checkpointing=use_gradient_checkpointing,
gradient_checkpointing_ratio=gradient_checkpointing_ratio,
use_flash_attn=use_flash_attn, use_temporal_causal=use_temporal_causal,
interp_condition_pos=interp_condition_pos, axes_dims_rope=[16, 24, 24],
)
elif model_name == "pyramid_mmdit":
dit = PyramidDiffusionMMDiT.from_pretrained(
model_path, torch_dtype=torch_dtype, use_gradient_checkpointing=use_gradient_checkpointing,
gradient_checkpointing_ratio=gradient_checkpointing_ratio,
use_flash_attn=use_flash_attn, use_t5_mask=True,
add_temp_pos_embed=True, temp_pos_embed_type='rope',
use_temporal_causal=use_temporal_causal, interp_condition_pos=interp_condition_pos,
)
else:
raise NotImplementedError(f"Unsupported DiT architecture, please set the model_name to `pyramid_flux` or `pyramid_mmdit`")
return dit
def build_text_encoder(
model_name : str,
model_path : str,
torch_dtype,
load_text_encoder: bool = True,
):
# The text encoder
if load_text_encoder:
if model_name == "pyramid_flux":
text_encoder = FluxTextEncoderWithMask(model_path, torch_dtype=torch_dtype)
elif model_name == "pyramid_mmdit":
text_encoder = SD3TextEncoderWithMask(model_path, torch_dtype=torch_dtype)
else:
raise NotImplementedError(f"Unsupported Text Encoder architecture, please set the model_name to `pyramid_flux` or `pyramid_mmdit`")
else:
text_encoder = None
return text_encoder
class PyramidDiTForVideoGeneration: class PyramidDiTForVideoGeneration:
""" """
The pyramid dit for both image and video generation, The running class wrapper The pyramid dit for both image and video generation, The running class wrapper
This class is mainly for fixed unit implementation: 1 + n + n + n This class is mainly for fixed unit implementation: 1 + n + n + n
""" """
def __init__(self, model_path, model_dtype, model_name, text_encoder_dtype, vae_dtype, use_gradient_checkpointing=False, return_log=True, def __init__(
model_variant="diffusion_transformer_768p", timestep_shift=1.0, stage_range=[0, 1/3, 2/3, 1], self,
sample_ratios=[1, 1, 1], scheduler_gamma=1/3, use_flash_attn=False, transformer,
load_text_encoder=True, load_vae=True, max_temporal_length=31, frame_per_unit=1, use_temporal_causal=True, model_dtype,
corrupt_ratio=1/3, interp_condition_pos=True, stages=[1, 2, 4], fp8_fastmode=False, **kwargs, model_name,
return_log=True,
timestep_shift=1.0,
stage_range=[0, 1/3, 2/3, 1],
sample_ratios=[1, 1, 1],
scheduler_gamma=1/3,
max_temporal_length=31,
frame_per_unit=1,
use_temporal_causal=True,
corrupt_ratio=1/3,
interp_condition_pos=True,
stages=[1, 2, 4],
**kwargs,
): ):
super().__init__() super().__init__()
self.dit = transformer
if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
self.dtype = torch.bfloat16 self.dtype = torch.bfloat16
else: else:
@@ -118,43 +62,6 @@ class PyramidDiTForVideoGeneration:
self.corrupt_ratio = corrupt_ratio self.corrupt_ratio = corrupt_ratio
self.model_name = model_name self.model_name = model_name
dit_path = os.path.join(model_path, model_variant)
# The dit
self.dit = build_pyramid_dit(
model_name, dit_path, self.dtype,
use_flash_attn=use_flash_attn,
interp_condition_pos=interp_condition_pos, use_gradient_checkpointing=use_gradient_checkpointing,
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 model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2] and fp8_fastmode:
from ..fp8_optimization import convert_fp8_linear
convert_fp8_linear(self.dit, torch.bfloat16)
# The text encoder
self.text_encoder = build_text_encoder(
model_name, model_path, text_encoder_dtype, load_text_encoder=load_text_encoder,
)
self.load_text_encoder = load_text_encoder
# The base video vae decoder
if load_vae:
self.vae = CausalVideoVAE.from_pretrained(os.path.join(model_path, 'causal_video_vae'), torch_dtype=vae_dtype, interpolate=False)
# Freeze vae
for parameter in self.vae.parameters():
parameter.requires_grad = False
else:
self.vae = None
# For the image latent # For the image latent
self.vae_shift_factor = 0.1490 self.vae_shift_factor = 0.1490
self.vae_scale_factor = 1 / 1.8415 self.vae_scale_factor = 1 / 1.8415
@@ -183,37 +90,7 @@ class PyramidDiTForVideoGeneration:
self.cfg_rate = 0.1 self.cfg_rate = 0.1
self.return_log = return_log self.return_log = return_log
self.use_flash_attn = use_flash_attn self.use_flash_attn = False
def load_checkpoint(self, checkpoint_path, model_key='model', **kwargs):
checkpoint = torch.load(checkpoint_path, map_location='cpu')
dit_checkpoint = OrderedDict()
for key in checkpoint:
if key.startswith('vae') or key.startswith('text_encoder'):
continue
if key.startswith('dit'):
new_key = key.split('.')
new_key = '.'.join(new_key[1:])
dit_checkpoint[new_key] = checkpoint[key]
else:
dit_checkpoint[key] = checkpoint[key]
load_result = self.dit.load_state_dict(dit_checkpoint, strict=True)
print(f"Load checkpoint from {checkpoint_path}, load result: {load_result}")
def load_vae_checkpoint(self, vae_checkpoint_path, model_key='model'):
checkpoint = torch.load(vae_checkpoint_path, map_location='cpu')
checkpoint = checkpoint[model_key]
loaded_checkpoint = OrderedDict()
for key in checkpoint.keys():
if key.startswith('vae.'):
new_key = key.split('.')
new_key = '.'.join(new_key[1:])
loaded_checkpoint[new_key] = checkpoint[key]
load_result = self.vae.load_state_dict(loaded_checkpoint)
print(f"Load the VAE from {vae_checkpoint_path}, load result: {load_result}")
@torch.no_grad() @torch.no_grad()
def get_pyramid_latent(self, x, stage_num): def get_pyramid_latent(self, x, stage_num):
+1 -1
View File
@@ -117,7 +117,7 @@ class CausalVideoVAE(ModelMixin, ConfigMixin):
): ):
super().__init__() super().__init__()
print(f"The latent dimmension channes is {encoder_out_channels}") #print(f"The latent dimension channes is {encoder_out_channels}")
# pass init params to Encoder # pass init params to Encoder
self.encoder = CausalVaeEncoder( self.encoder = CausalVaeEncoder(