12 Commits
Author SHA1 Message Date
kijai 426e365304 update examples 2024-11-15 15:28:14 +02:00
kijai 3b572e270b allow fp32 to work 2024-11-15 15:23:50 +02:00
kijai 50f5d1fb37 fix fp8 2024-11-14 01:44:57 +02:00
kijai 723537bb75 allow torch.compiling more stuff 2024-11-13 23:59:10 +02:00
kijai acb19bcd29 Update nodes.py 2024-11-13 18:47:29 +02:00
kijai 1d9fe8c695 Update nodes.py 2024-11-07 21:50:56 +02:00
kijai 3aeea32237 Add latent preview, cleanup 2024-11-01 14:38:28 +02:00
kijai bc8d4360fe Merge branch 'main' of https://github.com/kijai/ComfyUI-PyramidFlowWrapper 2024-10-30 16:08:08 +02:00
kijai ecd00cf903 cpu_offload 2024-10-30 16:08:00 +02:00
Jukka Seppänen eb5668b88b Update README.md 2024-10-30 15:53:19 +02:00
Jukka Seppänen 7387ea5de5 Update README.md 2024-10-30 15:52:46 +02:00
kijai 3e14b0d77e Big update, refactor model loading, breaks backwards compatibility
archiving old  version in legacy branch, probably not going to support SD3 version going forward
2024-10-30 15:48:32 +02:00
25 changed files with 2139 additions and 3387 deletions
+13 -64
View File
@@ -1,72 +1,21 @@
# ComfyUI wrapper nodes for [Pyramid-Flow](https://github.com/jy0205/Pyramid-Flow)
## UPDATE
As the first Flux version is out, I'm dropping the SD3 support and refactored the whole thing, if you still want to use the old nodes they are archived in the legacy branch
https://github.com/user-attachments/assets/9592bfed-9cf8-438a-b1e6-969d6b13ec12
The fluxmini version can run with 7GB VRAM, currently it only supports 5 second videos (temp 16).
Fp8 severely reduces quality and is not recommended, only use it if you must.
Download models from:
https://huggingface.co/Kijai/pyramid-flow-comfy/tree/main
To `ComfyUI/models/diffusion_models` and `ComfyUI/models/vae`
https://github.com/user-attachments/assets/1372549a-4b4e-4569-a062-8f72880e8c4e
todo
- refactor to use comfy text encoding instead
- optimize memory use
https://github.com/user-attachments/assets/d0bd38eb-6378-4cfa-ae55-1b4498b7ce84
Besides text encoder, which can peak at ~12GB VRAM use, this should run at 9-10GB VRAM when using 1280x768.
With fp8 and 384p model it can fit under 6GB too. Note that these tests were done on 4090, older cards may not support every optimization.
Resolutions outside the model defaults perform poorly.
Model loading has not been optimized at all yet, currently needs everything (choose either of the transformers) from here:
https://huggingface.co/rain1011/pyramid-flow-sd3/tree/main
to:
`ComfyUI/models/pyramidflow/pyramid-flow-sd3`
So that the directory structure is as follows:
```
\ComfyUI\models\pyramidflow\pyramid-flow-sd3
├───causal_video_vae
│ config.json
│ diffusion_pytorch_model.safetensors
│
├───diffusion_transformer_384p
│ config.json
│ diffusion_pytorch_model.safetensors
│
├───diffusion_transformer_768p
│ config.json
│ diffusion_pytorch_model.safetensors
│
├───text_encoder
│ config.json
│ model.safetensors
│
├───text_encoder_2
│ config.json
│ model.safetensors
│
├───text_encoder_3
│ config.json
│ model-00001-of-00002.safetensors
│ model-00002-of-00002.safetensors
│ model.safetensors.index.json
│
├───tokenizer
│ merges.txt
│ special_tokens_map.json
│ tokenizer_config.json
│ vocab.json
│
├───tokenizer_2
│ merges.txt
│ special_tokens_map.json
│ tokenizer_config.json
│ vocab.json
│
└───tokenizer_3
special_tokens_map.json
spiece.model
tokenizer.json
tokenizer_config.json
```
Original repo: https://github.com/jy0205/Pyramid-Flow
+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
}
@@ -0,0 +1,717 @@
{
"last_node_id": 59,
"last_link_id": 95,
"nodes": [
{
"id": 39,
"type": "Note",
"pos": {
"0": 30,
"1": 650
},
"size": {
"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": 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": 57,
"type": "ImageScale",
"pos": {
"0": 635,
"1": 833
},
"size": {
"0": 315,
"1": 130
},
"flags": {},
"order": 8,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 88
},
{
"name": "width",
"type": "INT",
"link": 91,
"widget": {
"name": "width"
}
},
{
"name": "height",
"type": "INT",
"link": 93,
"widget": {
"name": "height"
}
}
],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
89
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "ImageScale"
},
"widgets_values": [
"lanczos",
640,
384,
"center"
]
},
{
"id": 58,
"type": "PrimitiveNode",
"pos": {
"0": 347,
"1": 817
},
"size": {
"0": 256.92181396484375,
"1": 82
},
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [
91,
94
],
"slot_index": 0,
"widget": {
"name": "width"
}
}
],
"title": "width",
"properties": {
"Run widget replace on values": false
},
"widgets_values": [
640,
"fixed"
]
},
{
"id": 59,
"type": "PrimitiveNode",
"pos": {
"0": 350,
"1": 945
},
"size": {
"0": 251.2918701171875,
"1": 82
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "INT",
"type": "INT",
"links": [
93,
95
],
"slot_index": 0,
"widget": {
"name": "height"
}
}
],
"title": "height",
"properties": {
"Run widget replace on values": false
},
"widgets_values": [
384,
"fixed"
]
},
{
"id": 54,
"type": "PyramidFlowVAEEncode",
"pos": {
"0": 993,
"1": 786
},
"size": {
"0": 315,
"1": 102
},
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "vae",
"type": "PYRAMIDFLOWVAE",
"link": 82
},
{
"name": "image",
"type": "IMAGE",
"link": 89
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
83
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowVAEEncode"
},
"widgets_values": [
false,
0.25
]
},
{
"id": 55,
"type": "LoadImage",
"pos": {
"0": -32,
"1": 842
},
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
88
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"videoframe_812.png",
"image"
]
},
{
"id": 53,
"type": "PyramidFlowTextEncode",
"pos": {
"0": 444,
"1": 476
},
"size": {
"0": 437.19818115234375,
"1": 269.9795837402344
},
"flags": {},
"order": 7,
"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": [
"FPV flying over seaside cliffs while the sun is setting, 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": 50,
"type": "PyramidFlowSampler",
"pos": {
"0": 1046,
"1": 144
},
"size": {
"0": 315,
"1": 314
},
"flags": {},
"order": 10,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "PYRAMIDFLOWMODEL",
"link": 74
},
{
"name": "prompt_embeds",
"type": "PYRAMIDFLOWPROMPT",
"link": 81
},
{
"name": "input_latent",
"type": "LATENT",
"link": 83,
"shape": 7
},
{
"name": "width",
"type": "INT",
"link": 94,
"widget": {
"name": "width"
}
},
{
"name": "height",
"type": "INT",
"link": 95,
"widget": {
"name": "height"
}
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
76
]
}
],
"properties": {
"Node name for S&R": "PyramidFlowSampler"
},
"widgets_values": [
640,
384,
"20, 20, 20",
"10, 10, 10",
16,
7,
4,
44664248661402,
"fixed",
false
]
},
{
"id": 51,
"type": "VHS_VideoCombine",
"pos": {
"0": 1420,
"1": 113
},
"size": [
1018.2306518554688,
310
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 77
},
{
"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_00131.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 24
},
"muted": false
}
}
},
{
"id": 43,
"type": "PyramidFlowVAELoader",
"pos": {
"0": 250,
"1": 282
},
"size": {
"0": 411.12652587890625,
"1": 82
},
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"type": "PYRAMIDFLOW_COMPILEARGS",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "pyramidflow_vae",
"type": "PYRAMIDFLOWVAE",
"links": [
71,
82
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowVAELoader"
},
"widgets_values": [
"pyramidflow\\pyramid_flow_vae_bf16.safetensors",
"bf16"
]
},
{
"id": 48,
"type": "PyramidFlowVAEDecode",
"pos": {
"0": 1051,
"1": 534
},
"size": {
"0": 315,
"1": 150
},
"flags": {},
"order": 11,
"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,
0.25,
true,
true
]
},
{
"id": 40,
"type": "PyramidFlowTransformerLoader",
"pos": {
"0": 230,
"1": 114
},
"size": {
"0": 444.05462646484375,
"1": 106
},
"flags": {},
"order": 6,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"type": "PYRAMIDFLOW_COMPILEARGS",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "pyramidflow_model",
"type": "PYRAMIDFLOWMODEL",
"links": [
74
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowTransformerLoader"
},
"widgets_values": [
"pyramidflow\\pyramid_flow_miniflux_bf16_v2.safetensors",
"bf16",
false
]
}
],
"links": [
[
71,
43,
0,
48,
0,
"PYRAMIDFLOWVAE"
],
[
74,
40,
0,
50,
0,
"PYRAMIDFLOWMODEL"
],
[
76,
50,
0,
48,
1,
"LATENT"
],
[
77,
48,
0,
51,
0,
"IMAGE"
],
[
80,
37,
0,
53,
0,
"CLIP"
],
[
81,
53,
0,
50,
1,
"PYRAMIDFLOWPROMPT"
],
[
82,
43,
0,
54,
0,
"PYRAMIDFLOWVAE"
],
[
83,
54,
0,
50,
2,
"LATENT"
],
[
88,
55,
0,
57,
0,
"IMAGE"
],
[
89,
57,
0,
54,
1,
"IMAGE"
],
[
91,
58,
0,
57,
1,
"INT"
],
[
93,
59,
0,
57,
2,
"INT"
],
[
94,
58,
0,
50,
3,
"INT"
],
[
95,
59,
0,
50,
4,
"INT"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 0.6934334949442648,
"offset": [
667.6026332109438,
191.25659609525596
]
}
},
"version": 0.4
}
@@ -1,31 +1,88 @@
{
"last_node_id": 39,
"last_link_id": 54,
"last_node_id": 53,
"last_link_id": 81,
"nodes": [
{
"id": 9,
"type": "PyramidFlowSampler",
"id": 39,
"type": "Note",
"pos": {
"0": 1059,
"1": 497
"0": 30,
"1": 650
},
"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": 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
},
"flags": {},
"order": 4,
"order": 5,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "PYRAMIDFLOWMODEL",
"link": 7
"link": 74
},
{
"name": "prompt_embeds",
"type": "PYRAMIDFLOWPROMPT",
"link": 54
"link": 81
},
{
"name": "input_latent",
@@ -35,20 +92,12 @@
}
],
"outputs": [
{
"name": "model",
"type": "PYRAMIDFLOWMODEL",
"links": [
8
]
},
{
"name": "samples",
"type": "LATENT",
"links": [
9
],
"slot_index": 1
76
]
}
],
"properties": {
@@ -60,76 +109,112 @@
"20, 20, 20",
"10, 10, 10",
16,
9,
7,
5,
44664248661395,
44664248661402,
"fixed",
""
false
]
},
{
"id": 8,
"type": "PyramidFlowVAEDecode",
"id": 43,
"type": "PyramidFlowVAELoader",
"pos": {
"0": 1161,
"1": 873
"0": 250,
"1": 282
},
"size": {
"0": 315,
"1": 102
"0": 411.12652587890625,
"1": 82
},
"flags": {},
"order": 5,
"order": 2,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "PYRAMIDFLOWMODEL",
"link": 8
},
{
"name": "samples",
"type": "LATENT",
"link": 9
"name": "compile_args",
"type": "PYRAMIDFLOW_COMPILEARGS",
"link": null,
"shape": 7
}
],
"outputs": [
{
"name": "images",
"type": "IMAGE",
"name": "pyramidflow_vae",
"type": "PYRAMIDFLOWVAE",
"links": [
53
71
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "PyramidFlowVAEDecode"
"Node name for S&R": "PyramidFlowVAELoader"
},
"widgets_values": [
256,
2
"pyramidflow\\pyramid_flow_vae_bf16.safetensors",
"bf16"
]
},
{
"id": 14,
"id": 53,
"type": "PyramidFlowTextEncode",
"pos": {
"0": 444,
"1": 476
},
"size": {
"0": 437.19818115234375,
"1": 269.9795837402344
},
"flags": {},
"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": 1541,
"1": 339
"0": 1420,
"1": 113
},
"size": [
1698.6201171875,
1331.1720703125
1018.2306518554688,
310
],
"flags": {},
"order": 6,
"order": 7,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 53
"link": 77
},
{
"name": "audio",
@@ -174,7 +259,7 @@
"hidden": false,
"paused": false,
"params": {
"filename": "PyramidFlow_00089.mp4",
"filename": "PyramidFlow_00129.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
@@ -185,186 +270,139 @@
}
},
{
"id": 36,
"type": "PyramidFlowTextEncodeComfy",
"id": 40,
"type": "PyramidFlowTransformerLoader",
"pos": {
"0": 597,
"1": 779
"0": 225,
"1": 107
},
"size": {
"0": 400,
"1": 200
"0": 444.05462646484375,
"1": 106
},
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "CLIP",
"link": 47
"name": "compile_args",
"type": "PYRAMIDFLOW_COMPILEARGS",
"link": null,
"shape": 7
}
],
"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
74
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "DownloadAndLoadPyramidFlowModel"
"Node name for S&R": "PyramidFlowTransformerLoader"
},
"widgets_values": [
"rain1011/pyramid-flow-miniflux",
"diffusion_transformer_384p",
"bf16",
"bf16",
"pyramidflow\\pyramid_flow_miniflux_bf16_v2.safetensors",
"bf16",
false
]
},
{
"id": 48,
"type": "PyramidFlowVAEDecode",
"pos": {
"0": 1051,
"1": 534
},
"size": {
"0": 315,
"1": 150
},
"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,
0.25,
true,
true
]
}
],
"links": [
[
7,
5,
71,
43,
0,
9,
48,
0,
"PYRAMIDFLOWVAE"
],
[
74,
40,
0,
50,
0,
"PYRAMIDFLOWMODEL"
],
[
8,
9,
76,
50,
0,
8,
0,
"PYRAMIDFLOWMODEL"
],
[
9,
9,
1,
8,
48,
1,
"LATENT"
],
[
47,
37,
77,
48,
0,
36,
0,
"CLIP"
],
[
53,
8,
0,
14,
51,
0,
"IMAGE"
],
[
54,
36,
80,
37,
0,
9,
53,
0,
"CLIP"
],
[
81,
53,
0,
50,
1,
"PYRAMIDFLOWPROMPT"
]
@@ -373,10 +411,10 @@
"config": {},
"extra": {
"ds": {
"scale": 0.6303940863129696,
"scale": 0.6934334949442648,
"offset": [
274.1852517840429,
-178.2662230728557
667.6026332109438,
191.25659609525596
]
}
},
@@ -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
}
+6 -5
View File
@@ -36,10 +36,11 @@ def fp8_linear_forward(cls, original_dtype, input):
else:
return cls.original_forward(input)
def convert_fp8_linear(module, original_dtype):
def convert_fp8_linear(module, original_dtype, params_to_keep):
setattr(module, "fp8_matmul_enabled", True)
for name, module in module.named_modules():
if isinstance(module, nn.Linear):
original_forward = module.forward
setattr(module, "original_forward", original_forward)
setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input))
if not any(keyword in name for keyword in params_to_keep):
if isinstance(module, nn.Linear):
original_forward = module.forward
setattr(module, "original_forward", original_forward)
setattr(module, "forward", lambda input, m=module: fp8_linear_forward(m, original_dtype, input))
+87
View File
@@ -0,0 +1,87 @@
import io
import torch
from PIL import Image
import struct
import numpy as np
from comfy.cli_args import args, LatentPreviewMethod
from comfy.taesd.taesd import TAESD
import comfy.model_management
import folder_paths
import comfy.utils
import logging
MAX_PREVIEW_RESOLUTION = args.preview_size
def preview_to_image(latent_image):
latents_ubyte = (((latent_image + 1.0) / 2.0).clamp(0, 1) # change scale from -1..1 to 0..1
.mul(0xFF) # to 0..255
).to(device="cpu", dtype=torch.uint8, non_blocking=comfy.model_management.device_supports_non_blocking(latent_image.device))
return Image.fromarray(latents_ubyte.numpy())
class LatentPreviewer:
def decode_latent_to_preview(self, x0):
pass
def decode_latent_to_preview_image(self, preview_format, x0):
preview_image = self.decode_latent_to_preview(x0)
return ("GIF", preview_image, MAX_PREVIEW_RESOLUTION)
class Latent2RGBPreviewer(LatentPreviewer):
def __init__(self, latent_rgb_factors, latent_rgb_factors_bias=None):
latent_rgb_factors = [[0.05389399697934166, 0.025018778505575393, -0.009193515248318657], [0.02318250640590553, -0.026987363837713156, 0.040172639061236956], [0.046035451343323666, -0.02039565868920197, 0.01275569344290342], [-0.015559161155025095, 0.051403973219861246, 0.03179031307996347], [-0.02766167769640129, 0.03749545161530447, 0.003335141009473408], [0.05824598730479011, 0.021744367381243884, -0.01578925627951616], [0.05260929401500947, 0.0560165014956886, -0.027477296572565126], [0.018513891242931686, 0.041961785217662514, 0.004490763489747966], [0.024063060899760215, 0.065082853069653, 0.044343437673514896], [0.05250992323006226, 0.04361117432588933, 0.01030076055524387], [0.0038921710021782366, -0.025299228133723792, 0.019370764014574535], [-0.00011950534333568519, 0.06549370069727675, -0.03436712163379723], [-0.026020578032683626, -0.013341758571090847, -0.009119046570271953], [0.024412451175602937, 0.030135064560817174, -0.008355486384198006], [0.04002209845752687, -0.017341304390739463, 0.02818338690302971], [-0.032575108695213684, -0.009588338926775117, -0.03077312160940468]]
self.latent_rgb_factors = torch.tensor(latent_rgb_factors, device="cpu").transpose(0, 1)
self.latent_rgb_factors_bias = None
# if latent_rgb_factors_bias is not None:
# self.latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device="cpu")
def decode_latent_to_preview(self, x0):
self.latent_rgb_factors = self.latent_rgb_factors.to(dtype=x0.dtype, device=x0.device)
if self.latent_rgb_factors_bias is not None:
self.latent_rgb_factors_bias = self.latent_rgb_factors_bias.to(dtype=x0.dtype, device=x0.device)
latent_image = torch.nn.functional.linear(x0[0].permute(1, 2, 0), self.latent_rgb_factors,
bias=self.latent_rgb_factors_bias)
return preview_to_image(latent_image)
def get_previewer(device, latent_format):
previewer = None
method = args.preview_method
if method != LatentPreviewMethod.NoPreviews:
# TODO previewer methods
taesd_decoder_path = None
if latent_format.taesd_decoder_name is not None:
taesd_decoder_path = next(
(fn for fn in folder_paths.get_filename_list("vae_approx")
if fn.startswith(latent_format.taesd_decoder_name)),
""
)
taesd_decoder_path = folder_paths.get_full_path("vae_approx", taesd_decoder_path)
if method == LatentPreviewMethod.Auto:
method = LatentPreviewMethod.Latent2RGB
if previewer is None:
if latent_format.latent_rgb_factors is not None:
previewer = Latent2RGBPreviewer(latent_format.latent_rgb_factors, latent_format.latent_rgb_factors_bias)
return previewer
def prepare_callback(model, steps, x0_output_dict=None):
preview_format = "JPEG"
if preview_format not in ["JPEG", "PNG"]:
preview_format = "JPEG"
previewer = get_previewer(model.load_device, model.model.latent_format)
pbar = comfy.utils.ProgressBar(steps)
def callback(step, x0, x, total_steps):
if x0_output_dict is not None:
x0_output_dict["x0"] = x0
preview_bytes = None
if previewer:
preview_bytes = previewer.decode_latent_to_preview_image(preview_format, x0)
pbar.update_absolute(step + 1, total_steps, preview_bytes)
return callback
+327 -214
View File
@@ -1,6 +1,7 @@
import os
import torch
import folder_paths
import json
import comfy.model_management as mm
from comfy.utils import ProgressBar, load_torch_file
@@ -17,29 +18,118 @@ script_directory = os.path.dirname(os.path.abspath(__file__))
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"))
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
def INPUT_TYPES(s):
return {
"required": {
"model": (
[
"rain1011/pyramid-flow-sd3",
"rain1011/pyramid-flow-miniflux"
"backend": (["inductor","cudagraphs"], {"default": "inductor"}),
"fullgraph": ("BOOLEAN", {"default": False, "tooltip": "Enable full graph mode"}),
"mode": (["default", "max-autotune", "max-autotune-no-cudagraphs", "reduce-overhead"], {"default": "default"}),
"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)"}),
"dynamo_cache_size_limit": ("INT", {"default": 64, "min": 0, "max": 1024, "step": 1, "tooltip": "torch._dynamo.config.cache_size_limit"}),
},
}
RETURN_TYPES = ("PYRAMIDFLOW_COMPILEARGS",)
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"
],
),
"variant": (
["diffusion_transformer_384p", "diffusion_transformer_768p"],
),
def loadmodel(self, backend, fullgraph, mode, compile_whole_model, single_blocks, double_blocks, embedders, compile_rest, dynamo_cache_size_limit):
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,
"dynamo_cache_size_limit": dynamo_cache_size_limit,
}
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": {
"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", }),
"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"}),
"compile_args": ("PYRAMIDFLOW_COMPILEARGS", {"tooltip": "Optional torch.compile arguments",}),
}
}
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"}),
"enable_sequential_cpu_offload": ("BOOLEAN", {"default": False, "tooltip": "Enable sequential cpu offload, saves VRAM but is MUCH slower, do not use unless you have to"}),
},
"optional": {
"compile_args": ("PYRAMIDFLOW_COMPILEARGS", {"tooltip": "Optional torch.compile arguments",}),
}
}
@@ -48,112 +138,96 @@ class DownloadAndLoadPyramidFlowModel:
FUNCTION = "loadmodel"
CATEGORY = "PyramidFlowWrapper"
def loadmodel(self, model, variant, model_dtype, text_encoder_dtype, vae_dtype, fp8_fastmode):
def loadmodel(self, model, precision,enable_sequential_cpu_offload, compile_args=None):
device = mm.get_torch_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]
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]
model_path = folder_paths.get_full_path_or_raise("diffusion_models", model)
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])
variant_path = os.path.join(model_path, variant)
transformer_sd = load_torch_file(model_path)
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.extend["*text_encoder*", "*tokenizer*"]
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,
)
for key in transformer_sd:
if key.startswith("pos_embed."):
model_name = "pyramid_mmdit"
continue
else:
model_name = "pyramid_flux"
# # compilation
# if compile == "torch":
# torch._dynamo.config.suppress_errors = True
# pipe.transformer.to(memory_format=torch.channels_last)
# 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(
# pipe,
# 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 = {
"model": model,
"dtype": model_dtype,
"text_encoder_dtype": text_encoder_dtype,
"vae_dtype": vae_dtype,
}
return (pyramid_pipe,)
class CogVideoTextEncode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"clip": ("CLIP",),
"prompt": ("STRING", {"default": "", "multiline": True} ),
model_configs = {
"pyramid_flux": {
"config_file": "miniflux_transformer_config.json",
"transformer_class": PyramidFluxTransformer,
"params_to_keep": {"pos_embedding", "norm_k", "norm_q", "norm_v", "norm_added_k", "norm_added_q", "bias"}
},
"optional": {
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}),
"force_offload": ("BOOLEAN", {"default": True}),
"pyramid_mmdit": {
"config_file": "mmdit_transformer_config.json",
"transformer_class": PyramidDiffusionMMDiT,
"params_to_keep": {"pos_embedding"}
}
}
RETURN_TYPES = ("CONDITIONING",)
RETURN_NAMES = ("conditioning",)
FUNCTION = "process"
CATEGORY = "CogVideoWrapper"
config_info = model_configs[model_name]
config_path = os.path.join(script_directory, 'configs', config_info["config_file"])
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)
with open(config_path) as f:
config = json.load(f)
embeds = clip.encode_from_tokens(tokens, return_pooled=False, return_dict=False)
embeds *= strength
if force_offload:
clip.cond_stage_model.to(offload_device)
with (init_empty_weights() if is_accelerate_available else nullcontext()):
transformer = config_info["transformer_class"].from_config(config)
return (embeds, )
params_to_keep = config_info["params_to_keep"]
if is_accelerate_available:
logging.info("Using accelerate to load and assign model weights to device...")
base_dtype = torch.bfloat16 if dtype != torch.float32 else torch.float32
for name, param in transformer.named_parameters():
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
set_module_tensor_to_device(transformer, name, dtype=dtype_to_use, device=device, value=transformer_sd[name])
else:
transformer.load_state_dict(transformer_sd)
if dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
for param in transformer.parameters():
param.data = param.data.to(dtype)
if precision == "fp8_e4m3fn_fast":
from .fp8_optimization import convert_fp8_linear
convert_fp8_linear(transformer, torch.bfloat16, params_to_keep=params_to_keep)
transformer.to(device)
#torch.compile
if compile_args is not None:
torch._dynamo.config.force_parameter_static_shapes = False
torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"]
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"])
pyramid_model = PyramidDiTForVideoGeneration(transformer, dtype, model_name, device)
if enable_sequential_cpu_offload:
pyramid_model.enable_sequential_cpu_offload()
return (pyramid_model,)
#region Sampler
class PyramidFlowSampler:
@classmethod
def INPUT_TYPES(s):
@@ -177,8 +251,8 @@ class PyramidFlowSampler:
}
}
RETURN_TYPES = ("PYRAMIDFLOWMODEL", "LATENT", )
RETURN_NAMES = ("model","samples", )
RETURN_TYPES = ("LATENT", )
RETURN_NAMES = ("samples", )
FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper"
@@ -188,7 +262,16 @@ class PyramidFlowSampler:
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
dtype = model["dtype"]
if isinstance(model, dict):
pyramid_model = model["model"]
else:
pyramid_model = model
dtype = pyramid_model.dit.dtype
from .latent_preview import prepare_callback
callback = prepare_callback(model, temp)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
@@ -202,7 +285,7 @@ class PyramidFlowSampler:
if input_latent is None:
with autocast_context:
latents = model["model"].generate(
latents = pyramid_model.generate(
prompt_embeds_dict = prompt_embeds,
device=device,
num_inference_steps=first_frame_steps,
@@ -212,94 +295,35 @@ class PyramidFlowSampler:
temp=temp,
guidance_scale=guidance_scale, # The guidance for the first frame
video_guidance_scale=video_guidance_scale, # The guidance for the other video latent
output_type="latent",
callback=callback,
)
else:
with autocast_context:
latents = model["model"].generate_i2v(
latents = pyramid_model.generate_i2v(
prompt_embeds_dict = prompt_embeds,
input_image_latent=input_latent,
input_image_latent=input_latent["samples"],
device=device,
num_inference_steps=video_steps, #why's this a list
num_inference_steps=video_steps,
height=height,
width=width,
temp=temp,
video_guidance_scale=video_guidance_scale, # The guidance for the other video latent
output_type="latent",
callback=callback,
)
if not keep_model_loaded and not pyramid_model.sequential_offload_enabled:
pyramid_model.dit.to(offload_device)
if not keep_model_loaded:
model["model"].dit.to(offload_device)
return (model, {"samples": latents},)
return ({"samples": latents},)
#todo: sd3 version
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
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} ),
"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}),
}
}
@@ -324,7 +348,6 @@ class PyramidFlowTextEncodeComfy:
clip.cond_stage_model.t5_attention_mask = True
clip.cond_stage_model.to(device)#.to(torch.bfloat16)
clip.cond_stage_model.clip_l.to(device)
#positive
tokens = clip.tokenizer.t5xxl.tokenize_with_weights(positive_prompt, return_word_ids=False)
@@ -339,7 +362,7 @@ class PyramidFlowTextEncodeComfy:
if force_offload:
clip.cond_stage_model.to(offload_device)
clip.cond_stage_model.reset_clip_options()
embeds = {
"prompt_embeds": prompt_embeds.to(device),
"attention_mask": prompt_attention_mask["attention_mask"].to(device),
@@ -350,13 +373,17 @@ class PyramidFlowTextEncodeComfy:
}
return (embeds, )
#region VAE
class PyramidFlowVAEEncode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("PYRAMIDFLOWMODEL",),
"vae": ("PYRAMIDFLOWVAE",),
"image": ("IMAGE",),
"enable_tiling": ("BOOLEAN", {"default": False}),
"overlap_factor": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
},
}
@@ -365,43 +392,48 @@ class PyramidFlowVAEEncode:
FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper"
def sample(self, model, image):
def sample(self, vae, image, enable_tiling, overlap_factor):
B, H, W, C = image.shape
mm.soft_empty_cache()
self.vae = model["model"].vae
dtype = model["vae_dtype"]
dtype = vae.dtype
if enable_tiling:
vae.enable_tiling()
else:
vae.disable_tiling()
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
self.vae.disable_tiling()
vae.encode_tile_overlap_factor = overlap_factor
# For the image latent
self.vae_shift_factor = 0.1490
self.vae_scale_factor = 1 / 1.8415
vae_shift_factor = 0.1490
vae_scale_factor = 1 / 1.8415
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 = normalize(input_image_tensor)
input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1
input_image_tensor = normalize(input_image_tensor).unsqueeze(0)
input_image_tensor = rearrange(input_image_tensor, 'b t c h w -> b c t h w', t=B)
#input_image_tensor = input_image_tensor.unsqueeze(2) # Add temporal dimension t=1
input_image_tensor = input_image_tensor.to(dtype=dtype, device=device)
self.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]
self.vae.to(offload_device)
vae.to(device)
input_image_latent = (vae.encode(input_image_tensor).latent_dist.sample() - vae_shift_factor) * vae_scale_factor # [b c 1 h w]
vae.to(offload_device)
return (input_image_latent,)
return ({"samples": input_image_latent},)
class PyramidFlowVAEDecode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"model": ("PYRAMIDFLOWMODEL",),
"vae": ("PYRAMIDFLOWVAE",),
"samples": ("LATENT",),
"tile_sample_min_size": ("INT", {"default": 256, "min": 64, "max": 512, "step": 8}),
"overlap_factor": ("FLOAT", {"default": 0.25, "min": 0.0, "max": 1.0, "step": 0.01}),
"window_size": ("INT", {"default": 2, "min": 1, "max": 4, "step": 1}),
"enable_tiling": ("BOOLEAN", {"default": True}),
},
}
@@ -410,35 +442,39 @@ class PyramidFlowVAEDecode:
FUNCTION = "sample"
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, overlap_factor=0.25):
mm.soft_empty_cache()
latents = samples["samples"]
self.vae = model["model"].vae
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
self.vae.enable_tiling()
if enable_tiling:
vae.enable_tiling()
else:
vae.disable_tiling()
vae.decode_tile_overlap_factor = overlap_factor
# For the image latent
self.vae_shift_factor = 0.1490
self.vae_scale_factor = 1 / 1.8415
vae_shift_factor = 0.1490
vae_scale_factor = 1 / 1.8415
# For the video latent
self.vae_video_shift_factor = -0.2343
self.vae_video_scale_factor = 1 / 3.0986
vae_video_shift_factor = -0.2343
vae_video_scale_factor = 1 / 3.0986
self.vae.to(device)
latents = latents.to(self.vae.dtype)
vae.to(device)
latents = latents.to(vae.dtype)
if latents.shape[2] == 1:
latents = (latents / self.vae_scale_factor) + self.vae_shift_factor
latents = (latents / vae_scale_factor) + vae_shift_factor
else:
latents[:, :, :1] = (latents[:, :, :1] / self.vae_scale_factor) + self.vae_shift_factor
latents[:, :, 1:] = (latents[:, :, 1:] / self.vae_video_scale_factor) + self.vae_video_shift_factor
latents[:, :, :1] = (latents[:, :, :1] / vae_scale_factor) + vae_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 / 2 + 0.5).clamp(0, 1)
@@ -448,14 +484,88 @@ class PyramidFlowVAEDecode:
return (image,)
class PyramidFlowLatentPreview:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"samples": ("LATENT",),
# "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
# "min_val": ("FLOAT", {"default": -0.15, "min": -1.0, "max": 0.0, "step": 0.001}),
# "max_val": ("FLOAT", {"default": 0.15, "min": 0.0, "max": 1.0, "step": 0.001}),
},
}
RETURN_TYPES = ("IMAGE", "STRING", )
RETURN_NAMES = ("images", "latent_rgb_factors", )
FUNCTION = "sample"
CATEGORY = "PyramidFlowWrapper"
def sample(self, samples):#, seed, min_val, max_val):
mm.soft_empty_cache()
latents = samples["samples"].clone()
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
# For the image latent
vae_shift_factor = 0.1490
vae_scale_factor = 1 / 1.8415
# For the video latent
vae_video_shift_factor = -0.2343
vae_video_scale_factor = 1 / 3.0986
if latents.shape[2] == 1:
latents = (latents / vae_scale_factor) + vae_shift_factor
else:
latents[:, :, :1] = (latents[:, :, :1] / vae_scale_factor) + vae_shift_factor
latents[:, :, 1:] = (latents[:, :, 1:] / vae_video_scale_factor) + vae_video_shift_factor
latent_rgb_factors = [[0.05389399697934166, 0.025018778505575393, -0.009193515248318657], [0.02318250640590553, -0.026987363837713156, 0.040172639061236956], [0.046035451343323666, -0.02039565868920197, 0.01275569344290342], [-0.015559161155025095, 0.051403973219861246, 0.03179031307996347], [-0.02766167769640129, 0.03749545161530447, 0.003335141009473408], [0.05824598730479011, 0.021744367381243884, -0.01578925627951616], [0.05260929401500947, 0.0560165014956886, -0.027477296572565126], [0.018513891242931686, 0.041961785217662514, 0.004490763489747966], [0.024063060899760215, 0.065082853069653, 0.044343437673514896], [0.05250992323006226, 0.04361117432588933, 0.01030076055524387], [0.0038921710021782366, -0.025299228133723792, 0.019370764014574535], [-0.00011950534333568519, 0.06549370069727675, -0.03436712163379723], [-0.026020578032683626, -0.013341758571090847, -0.009119046570271953], [0.024412451175602937, 0.030135064560817174, -0.008355486384198006], [0.04002209845752687, -0.017341304390739463, 0.02818338690302971], [-0.032575108695213684, -0.009588338926775117, -0.03077312160940468]]
#import random
#random.seed(seed)
#latent_rgb_factors = [[random.uniform(min_val, max_val) for _ in range(3)] for _ in range(16)]
out_factors = latent_rgb_factors
print(latent_rgb_factors)
latent_rgb_factors_bias = [0,0,0]
latent_rgb_factors = torch.tensor(latent_rgb_factors, device=latents.device, dtype=latents.dtype).transpose(0, 1)
latent_rgb_factors_bias = torch.tensor(latent_rgb_factors_bias, device=latents.device, dtype=latents.dtype)
print("latent_rgb_factors", latent_rgb_factors.shape)
latent_images = []
for t in range(latents.shape[2]):
latent = latents[:, :, t, :, :]
latent = latent[0].permute(1, 2, 0)
latent_image = torch.nn.functional.linear(
latent,
latent_rgb_factors,
bias=latent_rgb_factors_bias
)
latent_images.append(latent_image)
latent_images = torch.stack(latent_images, dim=0)
print("latent_images", latent_images.shape)
latent_images_min = latent_images.min()
latent_images_max = latent_images.max()
latent_images = (latent_images - latent_images_min) / (latent_images_max - latent_images_min)
return (latent_images.float().cpu(), out_factors)
NODE_CLASS_MAPPINGS = {
"DownloadAndLoadPyramidFlowModel": DownloadAndLoadPyramidFlowModel,
"PyramidFlowSampler": PyramidFlowSampler,
"PyramidFlowVAEDecode": PyramidFlowVAEDecode,
"PyramidFlowTextEncode": PyramidFlowTextEncode,
"PyramidFlowVAEEncode": PyramidFlowVAEEncode,
"PyramidFlowTextEncodeComfy": PyramidFlowTextEncodeComfy,
"PyramidFlowTorchCompileSettings": PyramidFlowTorchCompileSettings,
"PyramidFlowTransformerLoader": PyramidFlowModelLoader,
"PyramidFlowVAELoader": PyramidFlowVAELoader,
"PyramidFlowLatentPreview": PyramidFlowLatentPreview
}
NODE_DISPLAY_NAME_MAPPINGS = {
@@ -464,5 +574,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"PyramidFlowVAEDecode" : "PyramidFlow VAE Decode",
"PyramidFlowTextEncode": "PyramidFlow Text Encode",
"PyramidFlowVAEEncode": "PyramidFlow VAE Encode",
"PyramidFlowTextEncodeComfy": "PyramidFlow Text Encode Comfy",
"PyramidFlowTorchCompileSettings": "PyramidFlow Torch Compile Settings",
"PyramidFlowTransformerLoader": "PyramidFlow Model Loader",
"PyramidFlowVAELoader": "PyramidFlow VAE Loader",
"PyramidFlowLatentPreview": "PyramidFlow Latent Preview"
}
-1
View File
@@ -1,3 +1,2 @@
from .modeling_pyramid_flux import PyramidFluxTransformer
from .modeling_text_encoder import FluxTextEncoderWithMask
from .modeling_flux_block import FluxSingleTransformerBlock, FluxTransformerBlock
+20 -306
View File
@@ -3,33 +3,31 @@ from typing import Any, Dict, List, Optional, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
import inspect
from einops import rearrange
from diffusers.utils import deprecate
from diffusers.models.activations import GEGLU, GELU, ApproximateGELU, SwiGLU
from .modeling_normalization import (
AdaLayerNormContinuous, AdaLayerNormZero,
AdaLayerNormZeroSingle, FP32LayerNorm, RMSNorm
)
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
all_to_all,
)
# try:
# from flash_attn.ops.triton.layer_norm import RMSNorm as FlashRMSNorm #slightly faster
# @torch.compiler.disable() #cause NaNs when compiled for some reason
# class RMSNorm(FlashRMSNorm):
# pass
# except:
from .modeling_normalization import RMSNorm
from .modeling_normalization import (AdaLayerNormZero, AdaLayerNormZeroSingle, FP32LayerNorm)
try:
from flash_attn import flash_attn_qkvpacked_func, flash_attn_func
from flash_attn.bert_padding import pad_input, unpad_input, index_first_axis
from flash_attn.flash_attn_interface import flash_attn_varlen_func
except:
flash_attn_func = None
flash_attn_qkvpacked_func = None
flash_attn_varlen_func = None
@torch.compiler.disable()
def compute_attention(query, key, value, attn_mask, dropout_p=0.0, is_causal=False):
return F.scaled_dot_product_attention(
query, key, value, dropout_p=dropout_p, is_causal=is_causal, attn_mask=attn_mask,
)
def apply_rope(xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
@@ -100,92 +98,6 @@ class FeedForward(nn.Module):
return hidden_states
class SequenceParallelVarlenFlashSelfAttentionWithT5Mask:
def __init__(self):
pass
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, encoder_attention_mask=None,
):
assert encoder_attention_mask is not None, "The encoder-hidden mask needed to be set"
batch_size = query.shape[0]
qkv_list = []
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
output_hidden = torch.zeros_like(qkv[:,:,0])
output_encoder_hidden = torch.zeros_like(encoder_qkv[:,:,0])
encoder_length = encoder_qkv.shape[1]
i_sum = 0
for i_p, length in enumerate(hidden_length):
# get the query, key, value from padding sequence
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, pad_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
indices = encoder_attention_mask[i_p]['indices']
qkv_list.append(index_first_axis(rearrange(concat_qkv_tokens, "b s ... -> (b s) ..."), indices))
i_sum += length
token_lengths = [x_.shape[0] for x_ in qkv_list]
qkv = torch.cat(qkv_list, dim=0)
query, key, value = qkv.unbind(1)
cu_seqlens = torch.cat([x_['seqlens_in_batch'] for x_ in encoder_attention_mask], dim=0)
max_seqlen_q = cu_seqlens.max().item()
max_seqlen_k = max_seqlen_q
cu_seqlens_q = F.pad(torch.cumsum(cu_seqlens, dim=0, dtype=torch.int32), (1, 0))
cu_seqlens_k = cu_seqlens_q.clone()
output = flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=0.0,
causal=False,
softmax_scale=scale,
)
# To merge the tokens
i_sum = 0;token_sum = 0
for i_p, length in enumerate(hidden_length):
tot_token_num = token_lengths[i_p]
stage_output = output[token_sum : token_sum + tot_token_num]
stage_output = pad_input(stage_output, encoder_attention_mask[i_p]['indices'], batch_size, encoder_length + length * sp_group_size)
stage_encoder_hidden_output = stage_output[:, :encoder_length]
stage_hidden_output = stage_output[:, encoder_length:]
stage_hidden_output = all_to_all(stage_hidden_output, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden[:, i_sum:i_sum+length] = stage_hidden_output
output_encoder_hidden[i_p::num_stages] = stage_encoder_hidden_output
token_sum += tot_token_num
i_sum += length
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden = output_hidden.flatten(2, 3)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
return output_hidden, output_encoder_hidden
class VarlenFlashSelfAttentionWithT5Mask:
def __init__(self):
@@ -262,69 +174,6 @@ class VarlenFlashSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenSelfAttentionWithT5Mask:
def __init__(self):
pass
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
encoder_length = encoder_qkv.shape[1]
i_sum = 0
output_encoder_hidden_list = []
output_hidden_list = []
for i_p, length in enumerate(hidden_length):
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, tot_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
query, key, value = concat_qkv_tokens.unbind(2) # [bs, tot_seq, nhead, dim]
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
stage_hidden_states = F.scaled_dot_product_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2) # [bs, tot_seq, nhead, dim]
output_encoder_hidden_list.append(stage_hidden_states[:, :encoder_length])
output_hidden = stage_hidden_states[:, encoder_length:]
output_hidden = all_to_all(output_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden_list.append(output_hidden)
i_sum += length
output_encoder_hidden = torch.stack(output_encoder_hidden_list, dim=1) # [b n s nhead d]
output_encoder_hidden = rearrange(output_encoder_hidden, 'b n s h d -> (b n) s h d')
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
output_hidden = torch.cat(output_hidden_list, dim=1).flatten(2, 3)
return output_hidden, output_encoder_hidden
class VarlenSelfAttentionWithT5Mask:
def __init__(self):
@@ -360,7 +209,7 @@ class VarlenSelfAttentionWithT5Mask:
value = value.transpose(1, 2)
# with torch.backends.cuda.sdp_kernel(enable_math=False, enable_flash=False, enable_mem_efficient=True):
stage_hidden_states = F.scaled_dot_product_attention(
stage_hidden_states = compute_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2).flatten(2, 3) # [bs, tot_seq, dim]
@@ -376,79 +225,6 @@ class VarlenSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenFlashAttnSingle:
def __init__(self):
pass
def __call__(
self, query, key, value, heads, scale,
hidden_length=None, image_rotary_emb=None, encoder_attention_mask=None,
):
assert encoder_attention_mask is not None, "The encoder-hidden mask needed to be set"
batch_size = query.shape[0]
qkv_list = []
num_stages = len(hidden_length)
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
output_hidden = torch.zeros_like(qkv[:,:,0])
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
i_sum = 0
for i_p, length in enumerate(hidden_length):
# get the query, key, value from padding sequence
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
if image_rotary_emb is not None:
qkv_tokens[:,:,0], qkv_tokens[:,:,1] = apply_rope(qkv_tokens[:,:,0], qkv_tokens[:,:,1], image_rotary_emb[i_p])
indices = encoder_attention_mask[i_p]['indices']
qkv_list.append(index_first_axis(rearrange(qkv_tokens, "b s ... -> (b s) ..."), indices))
i_sum += length
token_lengths = [x_.shape[0] for x_ in qkv_list]
qkv = torch.cat(qkv_list, dim=0)
query, key, value = qkv.unbind(1)
cu_seqlens = torch.cat([x_['seqlens_in_batch'] for x_ in encoder_attention_mask], dim=0)
max_seqlen_q = cu_seqlens.max().item()
max_seqlen_k = max_seqlen_q
cu_seqlens_q = F.pad(torch.cumsum(cu_seqlens, dim=0, dtype=torch.int32), (1, 0))
cu_seqlens_k = cu_seqlens_q.clone()
output = flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=0.0,
causal=False,
softmax_scale=scale,
)
# To merge the tokens
i_sum = 0;token_sum = 0
for i_p, length in enumerate(hidden_length):
tot_token_num = token_lengths[i_p]
stage_output = output[token_sum : token_sum + tot_token_num]
stage_output = pad_input(stage_output, encoder_attention_mask[i_p]['indices'], batch_size, length * sp_group_size)
stage_hidden_output = all_to_all(stage_output, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden[:, i_sum:i_sum+length] = stage_hidden_output
token_sum += tot_token_num
i_sum += length
output_hidden = output_hidden.flatten(2, 3)
return output_hidden
class VarlenFlashSelfAttnSingle:
def __init__(self):
@@ -515,56 +291,6 @@ class VarlenFlashSelfAttnSingle:
return output_hidden
class SequenceParallelVarlenAttnSingle:
def __init__(self):
pass
def __call__(
self, query, key, value, heads, scale,
hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
i_sum = 0
output_hidden_list = []
for i_p, length in enumerate(hidden_length):
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
if image_rotary_emb is not None:
qkv_tokens[:,:,0], qkv_tokens[:,:,1] = apply_rope(qkv_tokens[:,:,0], qkv_tokens[:,:,1], image_rotary_emb[i_p])
query, key, value = qkv_tokens.unbind(2) # [bs, tot_seq, nhead, dim]
query = query.transpose(1, 2).contiguous()
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
stage_hidden_states = F.scaled_dot_product_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2) # [bs, tot_seq, nhead, dim]
output_hidden = stage_hidden_states
output_hidden = all_to_all(output_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden_list.append(output_hidden)
i_sum += length
output_hidden = torch.cat(output_hidden_list, dim=1).flatten(2, 3)
return output_hidden
class VarlenSelfAttnSingle:
def __init__(self):
@@ -576,7 +302,6 @@ class VarlenSelfAttnSingle:
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
i_sum = 0
@@ -593,7 +318,7 @@ class VarlenSelfAttnSingle:
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
stage_hidden_states = F.scaled_dot_product_attention(
stage_hidden_states = compute_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2).flatten(2, 3) # [bs, tot_seq, dim]
@@ -732,15 +457,9 @@ class FluxSingleAttnProcessor2_0:
self.use_flash_attn = use_flash_attn
if self.use_flash_attn:
if is_sequence_parallel_initialized():
self.varlen_flash_attn = SequenceParallelVarlenFlashAttnSingle()
else:
self.varlen_flash_attn = VarlenFlashSelfAttnSingle()
self.varlen_flash_attn = VarlenFlashSelfAttnSingle()
else:
if is_sequence_parallel_initialized():
self.varlen_attn = SequenceParallelVarlenAttnSingle()
else:
self.varlen_attn = VarlenSelfAttnSingle()
self.varlen_attn = VarlenSelfAttnSingle()
def __call__(
self,
@@ -792,15 +511,9 @@ class FluxAttnProcessor2_0:
self.use_flash_attn = use_flash_attn
if self.use_flash_attn:
if is_sequence_parallel_initialized():
self.varlen_flash_attn = SequenceParallelVarlenFlashSelfAttentionWithT5Mask()
else:
self.varlen_flash_attn = VarlenFlashSelfAttentionWithT5Mask()
self.varlen_flash_attn = VarlenFlashSelfAttentionWithT5Mask()
else:
if is_sequence_parallel_initialized():
self.varlen_attn = SequenceParallelVarlenSelfAttentionWithT5Mask()
else:
self.varlen_attn = VarlenSelfAttentionWithT5Mask()
self.varlen_attn = VarlenSelfAttentionWithT5Mask()
def __call__(
self,
@@ -813,6 +526,7 @@ class FluxAttnProcessor2_0:
image_rotary_emb: Optional[torch.Tensor] = None,
) -> torch.FloatTensor:
# `sample` projections.
query = attn.to_q(hidden_states)
key = attn.to_k(hidden_states)
value = attn.to_v(hidden_states)
+19 -151
View File
@@ -1,30 +1,17 @@
from typing import Any, Dict, List, Optional, Union
import torch
import os
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from tqdm import tqdm
from diffusers.utils.torch_utils import randn_tensor
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.models.modeling_utils import ModelMixin
from diffusers.utils import is_torch_version
from .modeling_normalization import AdaLayerNormContinuous
from .modeling_embedding import CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings
from .modeling_embedding import CombinedTimestepTextProjEmbeddings
from .modeling_flux_block import FluxTransformerBlock, FluxSingleTransformerBlock
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
all_to_all,
)
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
assert dim % 2 == 0, "The dimension must be even."
@@ -269,13 +256,6 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
input_ids_list = [torch.cat([text_ids, image_ids], dim=1) for image_ids in image_ids_list]
image_rotary_emb = [self.pos_embed(input_ids) for input_ids in input_ids_list] # [bs, seq_len, 1, head_dim // 2, 2, 2]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
concat_output = True if self.training else False
image_rotary_emb = [all_to_all(x_.repeat(1, 1, sp_group_size, 1, 1, 1), sp_group, sp_group_size, scatter_dim=2, gather_dim=0, concat_output=concat_output) for x_ in image_rotary_emb]
input_ids_list = [all_to_all(input_ids.repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0, concat_output=concat_output) for input_ids in input_ids_list]
hidden_states, hidden_length = [], []
for sample_ in sample:
@@ -299,12 +279,6 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
pad_attention_mask = torch.ones((pad_batch_size, length), dtype=encoder_attention_mask.dtype).to(device)
pad_attention_mask = torch.cat([encoder_attention_mask[i_p::num_stages], pad_attention_mask], dim=1)
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
pad_attention_mask = all_to_all(pad_attention_mask.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0)
pad_attention_mask = pad_attention_mask.squeeze(2)
seqlens_in_batch = pad_attention_mask.sum(dim=-1, dtype=torch.int32)
indices = torch.nonzero(pad_attention_mask.flatten(), as_tuple=False).flatten()
@@ -331,13 +305,6 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
for i_p, length in enumerate(hidden_length):
image_ids_list.append(image_ids[i_p::num_stages][:, :length])
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
concat_output = True if self.training else False
text_ids = all_to_all(text_ids.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0, concat_output=concat_output).squeeze(2)
image_ids_list = [all_to_all(image_ids_.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0, concat_output=concat_output).squeeze(2) for image_ids_ in image_ids_list]
attention_mask = []
for i_p in range(len(hidden_length)):
image_ids = image_ids_list[i_p]
@@ -357,20 +324,11 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
output_hidden_list = []
batch_hidden_states = torch.split(batch_hidden_states, hidden_length, dim=1)
if is_sequence_parallel_initialized():
sp_group_size = get_sequence_parallel_world_size()
batch_size = batch_size // sp_group_size
for i_p, length in enumerate(hidden_length):
width, height, temp = widths[i_p], heights[i_p], temps[i_p]
trainable_token_num = trainable_token_list[i_p]
hidden_states = batch_hidden_states[i_p]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
hidden_states = all_to_all(hidden_states, sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
# only the trainable token are taking part in loss computation
hidden_states = hidden_states[:, -trainable_token_num:]
@@ -400,135 +358,45 @@ class PyramidFluxTransformer(ModelMixin, ConfigMixin):
hidden_states, hidden_length, temps, heights, widths, trainable_token_list, encoder_attention_mask, attention_mask, \
image_rotary_emb = self.merge_input(sample, encoder_hidden_length, encoder_attention_mask)
# split the long latents if necessary
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
concat_output = True if self.training else False
# sync the input hidden states
batch_hidden_states = []
for i_p, hidden_states_ in enumerate(hidden_states):
assert hidden_states_.shape[1] % sp_group_size == 0, "The sequence length should be divided by sequence parallel size"
hidden_states_ = all_to_all(hidden_states_, sp_group, sp_group_size, scatter_dim=1, gather_dim=0, concat_output=concat_output)
hidden_length[i_p] = hidden_length[i_p] // sp_group_size
batch_hidden_states.append(hidden_states_)
# sync the encoder hidden states
hidden_states = torch.cat(batch_hidden_states, dim=1)
encoder_hidden_states = all_to_all(encoder_hidden_states, sp_group, sp_group_size, scatter_dim=1, gather_dim=0, concat_output=concat_output)
temb = all_to_all(temb.unsqueeze(1).repeat(1, sp_group_size, 1), sp_group, sp_group_size, scatter_dim=1, gather_dim=0, concat_output=concat_output)
temb = temb.squeeze(1)
else:
hidden_states = torch.cat(hidden_states, dim=1)
hidden_states = torch.cat(hidden_states, dim=1)
for index_block, block in enumerate(self.transformer_blocks):
if self.training and self.gradient_checkpointing and (index_block <= int(len(self.transformer_blocks) * self.gradient_checkpointing_ratio)):
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
encoder_hidden_states, hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
attention_mask,
hidden_length,
image_rotary_emb,
**ckpt_kwargs,
)
else:
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
temb=temb,
attention_mask=attention_mask,
hidden_length=hidden_length,
image_rotary_emb=image_rotary_emb,
)
encoder_hidden_states, hidden_states = block(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
encoder_attention_mask=encoder_attention_mask,
temb=temb,
attention_mask=attention_mask,
hidden_length=hidden_length,
image_rotary_emb=image_rotary_emb,
)
# remerge for single attention block
num_stages = len(hidden_length)
batch_hidden_states = list(torch.split(hidden_states, hidden_length, dim=1))
concat_hidden_length = []
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_hidden_states = all_to_all(encoder_hidden_states, sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
for i_p in range(len(hidden_length)):
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
batch_hidden_states[i_p] = torch.cat([encoder_hidden_states[i_p::num_stages], batch_hidden_states[i_p]], dim=1)
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=1, gather_dim=0)
concat_hidden_length.append(batch_hidden_states[i_p].shape[1])
hidden_states = torch.cat(batch_hidden_states, dim=1)
for index_block, block in enumerate(self.single_transformer_blocks):
if self.training and self.gradient_checkpointing and (index_block <= int(len(self.single_transformer_blocks) * self.gradient_checkpointing_ratio)):
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
return custom_forward
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
hidden_states = torch.utils.checkpoint.checkpoint(
create_custom_forward(block),
hidden_states,
temb,
encoder_attention_mask,
attention_mask,
concat_hidden_length,
image_rotary_emb,
**ckpt_kwargs,
)
else:
hidden_states = block(
hidden_states=hidden_states,
temb=temb,
encoder_attention_mask=encoder_attention_mask, # used for
attention_mask=attention_mask,
hidden_length=concat_hidden_length,
image_rotary_emb=image_rotary_emb,
)
hidden_states = block(
hidden_states=hidden_states,
temb=temb,
encoder_attention_mask=encoder_attention_mask,
attention_mask=attention_mask,
hidden_length=concat_hidden_length,
image_rotary_emb=image_rotary_emb,
)
batch_hidden_states = list(torch.split(hidden_states, concat_hidden_length, dim=1))
for i_p in range(len(concat_hidden_length)):
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
batch_hidden_states[i_p] = batch_hidden_states[i_p][:, encoder_hidden_length :, ...]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
batch_hidden_states[i_p] = all_to_all(batch_hidden_states[i_p], sp_group, sp_group_size, scatter_dim=1, gather_dim=0)
hidden_states = torch.cat(batch_hidden_states, dim=1)
hidden_states = self.norm_out(hidden_states, temb, hidden_length=hidden_length)
hidden_states = self.proj_out(hidden_states)
@@ -1,141 +0,0 @@
import torch
import torch.nn as nn
import os
from transformers import (
CLIPTextModel,
CLIPTokenizer,
T5EncoderModel,
T5TokenizerFast,
)
from typing import Any, Callable, Dict, List, Optional, Union
class FluxTextEncoderWithMask(nn.Module):
def __init__(self, model_path, torch_dtype):
super().__init__()
# CLIP-G
self.tokenizer = CLIPTokenizer.from_pretrained(os.path.join(model_path, 'tokenizer'), torch_dtype=torch_dtype)
self.tokenizer_max_length = (
self.tokenizer.model_max_length if hasattr(self, "tokenizer") and self.tokenizer is not None else 77
)
self.text_encoder = CLIPTextModel.from_pretrained(os.path.join(model_path, 'text_encoder'), torch_dtype=torch_dtype)
# T5
self.tokenizer_2 = T5TokenizerFast.from_pretrained(os.path.join(model_path, 'tokenizer_2'))
self.text_encoder_2 = T5EncoderModel.from_pretrained(os.path.join(model_path, 'text_encoder_2'), torch_dtype=torch_dtype)
self._freeze()
def _freeze(self):
for param in self.parameters():
param.requires_grad = False
def _get_t5_prompt_embeds(
self,
prompt: Union[str, List[str]] = None,
num_images_per_prompt: int = 1,
max_sequence_length: int = 128,
device: Optional[torch.device] = None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer_2(
prompt,
padding="max_length",
max_length=max_sequence_length,
truncation=True,
return_length=False,
return_overflowing_tokens=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_attention_mask = text_inputs.attention_mask
prompt_attention_mask = prompt_attention_mask.to(device)
prompt_embeds = self.text_encoder_2(text_input_ids.to(device), attention_mask=prompt_attention_mask, output_hidden_states=False)[0]
dtype = self.text_encoder_2.dtype
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
_, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings and attention mask for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt, 1)
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, seq_len, -1)
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
prompt_attention_mask = prompt_attention_mask.repeat(num_images_per_prompt, 1)
return prompt_embeds, prompt_attention_mask
def _get_clip_prompt_embeds(
self,
prompt: Union[str, List[str]],
num_images_per_prompt: int = 1,
device: Optional[torch.device] = None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
text_inputs = self.tokenizer(
prompt,
padding="max_length",
max_length=self.tokenizer_max_length,
truncation=True,
return_overflowing_tokens=False,
return_length=False,
return_tensors="pt",
)
text_input_ids = text_inputs.input_ids
prompt_embeds = self.text_encoder(text_input_ids.to(device), output_hidden_states=False)
# Use pooled output of CLIPTextModel
prompt_embeds = prompt_embeds.pooler_output
prompt_embeds = prompt_embeds.to(dtype=self.text_encoder.dtype, device=device)
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_images_per_prompt)
prompt_embeds = prompt_embeds.view(batch_size * num_images_per_prompt, -1)
return prompt_embeds
def encode_prompt(self,
prompt,
num_images_per_prompt=1,
device=None,
):
prompt = [prompt] if isinstance(prompt, str) else prompt
batch_size = len(prompt)
pooled_prompt_embeds = self._get_clip_prompt_embeds(
prompt=prompt,
device=device,
num_images_per_prompt=num_images_per_prompt,
)
prompt_embeds, prompt_attention_mask = self._get_t5_prompt_embeds(
prompt=prompt,
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
def forward(self, input_prompts, device):
with torch.no_grad():
prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.encode_prompt(input_prompts, 1, device=device)
return prompt_embeds, prompt_attention_mask, pooled_prompt_embeds
-1
View File
@@ -1,2 +1 @@
from .modeling_pyramid_mmdit import PyramidDiffusionMMDiT
from .modeling_text_encoder import SD3TextEncoderWithMask
@@ -15,13 +15,6 @@ except:
flash_attn_varlen_func = None
print("Please install flash attention")
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
all_to_all,
)
from .modeling_normalization import AdaLayerNormZero, AdaLayerNormContinuous, RMSNorm
@@ -167,99 +160,6 @@ class VarlenFlashSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenFlashSelfAttentionWithT5Mask:
def __init__(self):
pass
def apply_rope(self, xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, encoder_attention_mask=None,
):
assert encoder_attention_mask is not None, "The encoder-hidden mask needed to be set"
batch_size = query.shape[0]
qkv_list = []
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
output_hidden = torch.zeros_like(qkv[:,:,0])
output_encoder_hidden = torch.zeros_like(encoder_qkv[:,:,0])
encoder_length = encoder_qkv.shape[1]
i_sum = 0
for i_p, length in enumerate(hidden_length):
# get the query, key, value from padding sequence
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, pad_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = self.apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
indices = encoder_attention_mask[i_p]['indices']
qkv_list.append(index_first_axis(rearrange(concat_qkv_tokens, "b s ... -> (b s) ..."), indices))
i_sum += length
token_lengths = [x_.shape[0] for x_ in qkv_list]
qkv = torch.cat(qkv_list, dim=0)
query, key, value = qkv.unbind(1)
cu_seqlens = torch.cat([x_['seqlens_in_batch'] for x_ in encoder_attention_mask], dim=0)
max_seqlen_q = cu_seqlens.max().item()
max_seqlen_k = max_seqlen_q
cu_seqlens_q = F.pad(torch.cumsum(cu_seqlens, dim=0, dtype=torch.int32), (1, 0))
cu_seqlens_k = cu_seqlens_q.clone()
output = flash_attn_varlen_func(
query,
key,
value,
cu_seqlens_q=cu_seqlens_q,
cu_seqlens_k=cu_seqlens_k,
max_seqlen_q=max_seqlen_q,
max_seqlen_k=max_seqlen_k,
dropout_p=0.0,
causal=False,
softmax_scale=scale,
)
# To merge the tokens
i_sum = 0;token_sum = 0
for i_p, length in enumerate(hidden_length):
tot_token_num = token_lengths[i_p]
stage_output = output[token_sum : token_sum + tot_token_num]
stage_output = pad_input(stage_output, encoder_attention_mask[i_p]['indices'], batch_size, encoder_length + length * sp_group_size)
stage_encoder_hidden_output = stage_output[:, :encoder_length]
stage_hidden_output = stage_output[:, encoder_length:]
stage_hidden_output = all_to_all(stage_hidden_output, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden[:, i_sum:i_sum+length] = stage_hidden_output
output_encoder_hidden[i_p::num_stages] = stage_encoder_hidden_output
token_sum += tot_token_num
i_sum += length
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden = output_hidden.flatten(2, 3)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
return output_hidden, output_encoder_hidden
class VarlenSelfAttentionWithT5Mask:
"""
@@ -321,79 +221,6 @@ class VarlenSelfAttentionWithT5Mask:
return output_hidden, output_encoder_hidden
class SequenceParallelVarlenSelfAttentionWithT5Mask:
"""
For chunk stage attention without using flash attention
"""
def __init__(self):
pass
def apply_rope(self, xq, xk, freqs_cis):
xq_ = xq.float().reshape(*xq.shape[:-1], -1, 1, 2)
xk_ = xk.float().reshape(*xk.shape[:-1], -1, 1, 2)
xq_out = freqs_cis[..., 0] * xq_[..., 0] + freqs_cis[..., 1] * xq_[..., 1]
xk_out = freqs_cis[..., 0] * xk_[..., 0] + freqs_cis[..., 1] * xk_[..., 1]
return xq_out.reshape(*xq.shape).type_as(xq), xk_out.reshape(*xk.shape).type_as(xk)
def __call__(
self, query, key, value, encoder_query, encoder_key, encoder_value,
heads, scale, hidden_length=None, image_rotary_emb=None, attention_mask=None,
):
assert attention_mask is not None, "The attention mask needed to be set"
num_stages = len(hidden_length)
encoder_qkv = torch.stack([encoder_query, encoder_key, encoder_value], dim=2) # [bs, sub_seq, 3, head, head_dim]
qkv = torch.stack([query, key, value], dim=2) # [bs, sub_seq, 3, head, head_dim]
# To sync the encoder query, key and values
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
encoder_qkv = all_to_all(encoder_qkv, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
encoder_length = encoder_qkv.shape[1]
i_sum = 0
output_encoder_hidden_list = []
output_hidden_list = []
for i_p, length in enumerate(hidden_length):
encoder_qkv_tokens = encoder_qkv[i_p::num_stages]
qkv_tokens = qkv[:, i_sum:i_sum+length]
qkv_tokens = all_to_all(qkv_tokens, sp_group, sp_group_size, scatter_dim=3, gather_dim=1) # [bs, seq, 3, sub_head, head_dim]
concat_qkv_tokens = torch.cat([encoder_qkv_tokens, qkv_tokens], dim=1) # [bs, tot_seq, 3, nhead, dim]
if image_rotary_emb is not None:
concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1] = self.apply_rope(concat_qkv_tokens[:,:,0], concat_qkv_tokens[:,:,1], image_rotary_emb[i_p])
query, key, value = concat_qkv_tokens.unbind(2) # [bs, tot_seq, nhead, dim]
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
stage_hidden_states = F.scaled_dot_product_attention(
query, key, value, dropout_p=0.0, is_causal=False, attn_mask=attention_mask[i_p],
)
stage_hidden_states = stage_hidden_states.transpose(1, 2) # [bs, tot_seq, nhead, dim]
output_encoder_hidden_list.append(stage_hidden_states[:, :encoder_length])
output_hidden = stage_hidden_states[:, encoder_length:]
output_hidden = all_to_all(output_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_hidden_list.append(output_hidden)
i_sum += length
output_encoder_hidden = torch.stack(output_encoder_hidden_list, dim=1) # [b n s nhead d]
output_encoder_hidden = rearrange(output_encoder_hidden, 'b n s h d -> (b n) s h d')
output_encoder_hidden = all_to_all(output_encoder_hidden, sp_group, sp_group_size, scatter_dim=1, gather_dim=2)
output_encoder_hidden = output_encoder_hidden.flatten(2, 3)
output_hidden = torch.cat(output_hidden_list, dim=1).flatten(2, 3)
return output_hidden, output_encoder_hidden
class JointAttention(nn.Module):
def __init__(
@@ -476,14 +303,8 @@ class JointAttention(nn.Module):
# print(f"Using flash-attention: {self.use_flash_attn}")
if self.use_flash_attn:
#if is_sequence_parallel_initialized():
# self.var_flash_attn = SequenceParallelVarlenFlashSelfAttentionWithT5Mask()
#else:
self.var_flash_attn = VarlenFlashSelfAttentionWithT5Mask()
else:
#if is_sequence_parallel_initialized():
#self.var_len_attn = SequenceParallelVarlenSelfAttentionWithT5Mask()
#else:
self.var_len_attn = VarlenSelfAttentionWithT5Mask()
@@ -13,16 +13,6 @@ from .modeling_embedding import PatchEmbed3D, CombinedTimestepConditionEmbedding
from .modeling_normalization import AdaLayerNormContinuous
from .modeling_mmdit_block import JointTransformerBlock
from ...trainer_misc import (
is_sequence_parallel_initialized,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
all_to_all,
)
#from IPython import embed
def rope(pos: torch.Tensor, dim: int, theta: int) -> torch.Tensor:
assert dim % 2 == 0, "The dimension must be even."
@@ -316,12 +306,6 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin):
pad_attention_mask = torch.ones((pad_batch_size, length), dtype=encoder_attention_mask.dtype).to(device)
pad_attention_mask = torch.cat([encoder_attention_mask[i_p::num_stages], pad_attention_mask], dim=1)
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
pad_attention_mask = all_to_all(pad_attention_mask.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0)
pad_attention_mask = pad_attention_mask.squeeze(2)
seqlens_in_batch = pad_attention_mask.sum(dim=-1, dtype=torch.int32)
indices = torch.nonzero(pad_attention_mask.flatten(), as_tuple=False).flatten()
@@ -347,12 +331,6 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin):
for i_p, length in enumerate(hidden_length):
image_ids_list.append(image_ids[i_p::num_stages][:, :length])
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
text_ids = all_to_all(text_ids.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0).squeeze(2)
image_ids_list = [all_to_all(image_ids_.unsqueeze(2).repeat(1, 1, sp_group_size), sp_group, sp_group_size, scatter_dim=2, gather_dim=0).squeeze(2) for image_ids_ in image_ids_list]
attention_mask = []
for i_p in range(len(hidden_length)):
image_ids = image_ids_list[i_p]
@@ -372,20 +350,11 @@ class PyramidDiffusionMMDiT(ModelMixin, ConfigMixin):
output_hidden_list = []
batch_hidden_states = torch.split(batch_hidden_states, hidden_length, dim=1)
if is_sequence_parallel_initialized():
sp_group_size = get_sequence_parallel_world_size()
batch_size = batch_size // sp_group_size
for i_p, length in enumerate(hidden_length):
width, height, temp = widths[i_p], heights[i_p], temps[i_p]
trainable_token_num = trainable_token_list[i_p]
hidden_states = batch_hidden_states[i_p]
if is_sequence_parallel_initialized():
sp_group = get_sequence_parallel_group()
sp_group_size = get_sequence_parallel_world_size()
hidden_states = all_to_all(hidden_states, sp_group, sp_group_size, scatter_dim=0, gather_dim=1)
# only the trainable token are taking part in loss computation
hidden_states = hidden_states[:, -trainable_token_num:]
+67 -221
View File
@@ -1,29 +1,17 @@
import torch
import os
from types import SimpleNamespace
import torch
import torch.nn.functional as F
from collections import OrderedDict
from einops import rearrange
from diffusers.utils.torch_utils import randn_tensor
import math
from tqdm import tqdm
from typing import List, Optional, Union
from typing import List, Optional, Union, Callable
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 accelerate import cpu_offload
from comfy.utils import ProgressBar
def compute_density_for_timestep_sampling(
@@ -41,73 +29,40 @@ def compute_density_for_timestep_sampling(
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:
"""
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
"""
def __init__(self, model_path, model_dtype, model_name, text_encoder_dtype, vae_dtype, use_gradient_checkpointing=False, return_log=True,
model_variant="diffusion_transformer_768p", timestep_shift=1.0, stage_range=[0, 1/3, 2/3, 1],
sample_ratios=[1, 1, 1], scheduler_gamma=1/3, use_flash_attn=False,
load_text_encoder=True, load_vae=True, max_temporal_length=31, frame_per_unit=1, use_temporal_causal=True,
corrupt_ratio=1/3, interp_condition_pos=True, stages=[1, 2, 4], fp8_fastmode=False, **kwargs,
):
def __init__(
self,
transformer,
model_dtype,
model_name,
main_device,
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__()
self.dit = transformer
self.device = main_device
self.sequential_offload_enabled = False
from comfy import latent_formats
self.model = SimpleNamespace(latent_format=latent_formats.Flux())
self.load_device = main_device
if model_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]:
self.dtype = torch.bfloat16
else:
@@ -118,43 +73,6 @@ class PyramidDiTForVideoGeneration:
self.corrupt_ratio = corrupt_ratio
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
self.vae_shift_factor = 0.1490
self.vae_scale_factor = 1 / 1.8415
@@ -178,42 +96,13 @@ class PyramidDiTForVideoGeneration:
)
#round the gamma as 1/3 seems to have issues on some systems
self.gamma = round(self.scheduler.config.gamma, 5)
self.dist = torch.distributions.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + self.gamma) - torch.ones(4, 4) * self.gamma)
print(f"The start sigmas and end sigmas of each stage is Start: {self.scheduler.start_sigmas}, End: {self.scheduler.end_sigmas}, Ori_start: {self.scheduler.ori_start_sigmas}")
self.cfg_rate = 0.1
self.return_log = return_log
self.use_flash_attn = use_flash_attn
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}")
self.use_flash_attn = False
@torch.no_grad()
def get_pyramid_latent(self, x, stage_num):
@@ -233,6 +122,14 @@ class PyramidDiTForVideoGeneration:
vae_latent_list = list(reversed(vae_latent_list))
return vae_latent_list
def _enable_sequential_cpu_offload(self, model, device):
self.sequential_offload_enabled = True
offload_buffers = len(model._parameters) > 0
cpu_offload(model, device, offload_buffers=offload_buffers)
def enable_sequential_cpu_offload(self):
self._enable_sequential_cpu_offload(self.dit, device=self.device)
def prepare_latents(
self,
batch_size,
@@ -255,12 +152,13 @@ class PyramidDiTForVideoGeneration:
return latents
def sample_block_noise(self, bs, ch, temp, height, width):
dist = torch.distributions.multivariate_normal.MultivariateNormal(torch.zeros(4), torch.eye(4) * (1 + self.gamma) - torch.ones(4, 4) * self.gamma)
block_number = bs * ch * temp * (height // 2) * (width // 2)
noise = torch.stack([dist.sample() for _ in range(block_number)]) # [block number, 4]
noise = torch.stack([self.dist.sample() for _ in range(block_number)]) # [block number, 4]
noise = rearrange(noise, '(b c t h w) (p q) -> b c t (h p) (w q)',b=bs,c=ch,t=temp,h=height//2,w=width//2,p=2,q=2)
return noise
@torch.no_grad()
def generate_one_unit(
self,
@@ -277,6 +175,7 @@ class PyramidDiTForVideoGeneration:
dtype,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
is_first_frame: bool = False,
callback = None,
):
stages = self.stages
intermed_latents = []
@@ -309,7 +208,6 @@ class PyramidDiTForVideoGeneration:
timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype)
latent_model_input = past_conditions[i_s] + [latent_model_input]
noise_pred = self.dit(
sample=[latent_model_input],
timestep_ratio=timestep,
@@ -320,10 +218,6 @@ class PyramidDiTForVideoGeneration:
noise_pred = noise_pred[0]
# nan_mask = torch.isnan(noise_pred)
# if torch.any(nan_mask):
# raise ValueError("nan in hidden_states")
# perform guidance
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
@@ -339,12 +233,10 @@ class PyramidDiTForVideoGeneration:
sample=latents,
generator=generator,
).prev_sample
#nan_mask = torch.isnan(latents)
#if torch.any(nan_mask):
# raise ValueError("nan in latents")
intermed_latents.append(latents)
return intermed_latents
@torch.no_grad()
@@ -364,32 +256,18 @@ class PyramidDiTForVideoGeneration:
alpha: float = 0.5,
num_images_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
output_type: Optional[str] = "pil",
callback: Optional[Callable] = None,
):
#device = self.device
dtype = self.dtype
assert temp % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
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
# else:
# assert isinstance(prompt, list)
# batch_size = len(prompt)
# prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt]
if isinstance(num_inference_steps, int):
num_inference_steps = [num_inference_steps] * len(self.stages)
elif isinstance(num_inference_steps, list) and len(num_inference_steps) < len(self.stages):
num_inference_steps = (num_inference_steps * len(self.stages))[:len(self.stages)]
# negative_prompt = negative_prompt or ""
# # Get the text embeddings
# prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device)
# negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device)
if use_linear_guidance:
max_guidance_scale = guidance_scale
guidance_scale_list = [max(max_guidance_scale - alpha * t_, min_guidance_scale) for t_ in range(temp+1)]
@@ -411,11 +289,6 @@ class PyramidDiTForVideoGeneration:
pooled_prompt_embeds = torch.cat([negative_pooled_prompt_embeds, positive_pooled_prompt_embeds], dim=0)
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, positive_prompt_attention_mask], dim=0)
# prompt_embeds = prompt_embeds.to(dtype)
# pooled_prompt_embeds = pooled_prompt_embeds.to(dtype)
# prompt_attention_mask = prompt_attention_mask.to(dtype)
# Create the initial random noise
num_channels_latents = (self.dit.config.in_channels // 4) if self.model_name == "pyramid_flux" else self.dit.config.in_channels
latents = self.prepare_latents(
@@ -442,18 +315,11 @@ class PyramidDiTForVideoGeneration:
num_units = temp // self.frame_per_unit
stages = self.stages
# # encode the image latents
# image_transform = transforms.Compose([
# transforms.ToTensor(),
# transforms.Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),
# ])
#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
#last_generated_latents = input_image_latent
self.dit.to(device)
if not self.sequential_offload_enabled:
self.dit.to(device)
comfy_pbar = ProgressBar(num_units)
for unit_index in tqdm(range(1, num_units + 1)):
@@ -504,20 +370,19 @@ class PyramidDiTForVideoGeneration:
dtype,
generator,
is_first_frame=False,
callback=callback
)
comfy_pbar.update(1)
if callback is not None:
callback(unit_index, intermed_latents[-1].detach()[0].permute(1,0,2,3), None, temp)
else:
comfy_pbar.update(1)
generated_latents_list.append(intermed_latents[-1])
#last_generated_latents = intermed_latents
generated_latents = torch.cat(generated_latents_list, dim=2)
if output_type == "latent":
image = generated_latents
else:
image = self.decode_latent(generated_latents)
return image
return generated_latents
@torch.no_grad()
def generate(
@@ -535,19 +400,11 @@ class PyramidDiTForVideoGeneration:
alpha: float = 0.5,
num_images_per_prompt: Optional[int] = 1,
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
output_type: Optional[str] = "pil",
device: Optional[torch.device] = None,
callback = None
):
assert (temp - 1) % self.frame_per_unit == 0, "The frames should be divided by frame_per unit"
# if isinstance(prompt, str):
# batch_size = 1
# prompt = prompt + ", hyper quality, Ultra HD, 8K" # adding this prompt to improve aesthetics
# else:
# assert isinstance(prompt, list)
# batch_size = len(prompt)
# prompt = [_ + ", hyper quality, Ultra HD, 8K" for _ in prompt]
if isinstance(num_inference_steps, int):
num_inference_steps = [num_inference_steps] * len(self.stages)
elif isinstance(num_inference_steps, list) and len(num_inference_steps) < len(self.stages):
@@ -558,19 +415,11 @@ class PyramidDiTForVideoGeneration:
elif isinstance(video_num_inference_steps, list) and len(video_num_inference_steps) < len(self.stages):
video_num_inference_steps = (video_num_inference_steps * len(self.stages))[:len(self.stages)]
#negative_prompt = negative_prompt or ""
# # Get the text embeddings
# self.text_encoder.to(device)
# prompt_embeds, prompt_attention_mask, pooled_prompt_embeds = self.text_encoder(prompt, device)
# negative_prompt_embeds, negative_prompt_attention_mask, negative_pooled_prompt_embeds = self.text_encoder(negative_prompt, device)
# self.text_encoder.to('cpu')
batch_size = prompt_embeds_dict['prompt_embeds'].shape[0]
if use_linear_guidance:
max_guidance_scale = guidance_scale
# guidance_scale_list = torch.linspace(max_guidance_scale, min_guidance_scale, temp).tolist()
guidance_scale_list = [max(max_guidance_scale - alpha * t_, min_guidance_scale) for t_ in range(temp)]
print(guidance_scale_list)
@@ -622,9 +471,9 @@ class PyramidDiTForVideoGeneration:
stages = self.stages
generated_latents_list = [] # The generated results
#last_generated_latents = None
self.dit.to(device)
if not self.sequential_offload_enabled:
self.dit.to(device)
comfy_pbar = ProgressBar(num_units)
for unit_index in tqdm(range(num_units)):
@@ -648,6 +497,7 @@ class PyramidDiTForVideoGeneration:
self.dtype,
generator,
is_first_frame=True,
callback=callback
)
else:
# prepare the condition latents
@@ -693,23 +543,19 @@ class PyramidDiTForVideoGeneration:
self.dtype,
generator,
is_first_frame=False,
callback=callback
)
comfy_pbar.update(1)
generated_latents_list.append(intermed_latents[-1])
#last_generated_latents = intermed_latents
if callback is not None:
callback(unit_index, intermed_latents[-1].detach()[0].permute(1,0,2,3), None, temp)
else:
comfy_pbar.update(1)
generated_latents = torch.cat(generated_latents_list, dim=2)
if output_type == "latent":
image = generated_latents
else:
image = self.decode_latent(generated_latents, device)
return image
@property
def device(self):
return next(self.dit.parameters()).device
return generated_latents
@property
def guidance_scale(self):
-25
View File
@@ -1,25 +0,0 @@
from .utils import (
create_optimizer,
get_rank,
get_world_size,
is_main_process,
is_dist_avail_and_initialized,
init_distributed_mode,
setup_for_distributed,
cosine_scheduler,
constant_scheduler,
)
from .sp_utils import (
is_sequence_parallel_initialized,
init_sequence_parallel_group,
get_sequence_parallel_group,
get_sequence_parallel_world_size,
get_sequence_parallel_rank,
get_sequence_parallel_group_rank,
get_sequence_parallel_proc_num,
init_sync_input_group,
get_sync_input_group,
)
from .communicate import all_to_all
-56
View File
@@ -1,56 +0,0 @@
import torch
import torch.distributed as dist
def _all_to_all(
input_: torch.Tensor,
world_size: int,
group: dist.ProcessGroup,
scatter_dim: int,
gather_dim: int,
):
if world_size == 1:
return input_
input_list = [t.contiguous() for t in torch.tensor_split(input_, world_size, scatter_dim)]
output_list = [torch.empty_like(input_list[0]) for _ in range(world_size)]
dist.all_to_all(output_list, input_list, group=group)
return torch.cat(output_list, dim=gather_dim).contiguous()
class _AllToAll(torch.autograd.Function):
@staticmethod
def forward(ctx, input_, process_group, world_size, scatter_dim, gather_dim):
ctx.process_group = process_group
ctx.scatter_dim = scatter_dim
ctx.gather_dim = gather_dim
ctx.world_size = world_size
output = _all_to_all(input_, ctx.world_size, process_group, scatter_dim, gather_dim)
return output
@staticmethod
def backward(ctx, grad_output):
grad_output = _all_to_all(
grad_output,
ctx.world_size,
ctx.process_group,
ctx.gather_dim,
ctx.scatter_dim,
)
return (
grad_output,
None,
None,
None,
None,
)
def all_to_all(
input_: torch.Tensor,
process_group: dist.ProcessGroup,
world_size: int = 1,
scatter_dim: int = 2,
gather_dim: int = 1,
):
return _AllToAll.apply(input_, process_group, world_size, scatter_dim, gather_dim)
-97
View File
@@ -1,97 +0,0 @@
import os
import torch
from .utils import is_dist_avail_and_initialized, get_rank
SEQ_PARALLEL_GROUP = None
SEQ_PARALLEL_SIZE = None
SEQ_PARALLEL_PROC_NUM = None # using how many process for sequence parallel
SYNC_INPUT_GROUP = None
SYNC_INPUT_SIZE = None
def is_sequence_parallel_initialized():
if SEQ_PARALLEL_GROUP is None:
return False
else:
return True
def init_sequence_parallel_group(args):
global SEQ_PARALLEL_GROUP
global SEQ_PARALLEL_SIZE
global SEQ_PARALLEL_PROC_NUM
assert SEQ_PARALLEL_GROUP is None, "sequence parallel group is already initialized"
assert is_dist_avail_and_initialized(), "The pytorch distributed should be initialized"
SEQ_PARALLEL_SIZE = args.sp_group_size
print(f"Setting the Sequence Parallel Size {SEQ_PARALLEL_SIZE}")
rank = torch.distributed.get_rank()
world_size = torch.distributed.get_world_size()
if args.sp_proc_num == -1:
SEQ_PARALLEL_PROC_NUM = world_size
else:
SEQ_PARALLEL_PROC_NUM = args.sp_proc_num
assert SEQ_PARALLEL_PROC_NUM % SEQ_PARALLEL_SIZE == 0, "The process needs to be evenly divided"
for i in range(0, SEQ_PARALLEL_PROC_NUM, SEQ_PARALLEL_SIZE):
ranks = list(range(i, i + SEQ_PARALLEL_SIZE))
group = torch.distributed.new_group(ranks)
if rank in ranks:
SEQ_PARALLEL_GROUP = group
break
def init_sync_input_group(args):
global SYNC_INPUT_GROUP
global SYNC_INPUT_SIZE
assert SYNC_INPUT_GROUP is None, "parallel group is already initialized"
assert is_dist_avail_and_initialized(), "The pytorch distributed should be initialized"
SYNC_INPUT_SIZE = args.max_frames
rank = torch.distributed.get_rank()
world_size = torch.distributed.get_world_size()
for i in range(0, world_size, SYNC_INPUT_SIZE):
ranks = list(range(i, i + SYNC_INPUT_SIZE))
group = torch.distributed.new_group(ranks)
if rank in ranks:
SYNC_INPUT_GROUP = group
break
def get_sequence_parallel_group():
assert SEQ_PARALLEL_GROUP is not None, "sequence parallel group is not initialized"
return SEQ_PARALLEL_GROUP
def get_sync_input_group():
return SYNC_INPUT_GROUP
def get_sequence_parallel_world_size():
assert SEQ_PARALLEL_SIZE is not None, "sequence parallel size is not initialized"
return SEQ_PARALLEL_SIZE
def get_sequence_parallel_rank():
assert SEQ_PARALLEL_SIZE is not None, "sequence parallel size is not initialized"
rank = get_rank()
cp_rank = rank % SEQ_PARALLEL_SIZE
return cp_rank
def get_sequence_parallel_group_rank():
assert SEQ_PARALLEL_SIZE is not None, "sequence parallel size is not initialized"
rank = get_rank()
cp_group_rank = rank // SEQ_PARALLEL_SIZE
return cp_group_rank
def get_sequence_parallel_proc_num():
return SEQ_PARALLEL_PROC_NUM
-377
View File
@@ -1,377 +0,0 @@
import os
import math
import time
import json
from collections import defaultdict, deque
import datetime
import numpy as np
import torch
from torch import optim as optim
import torch.distributed as dist
#from tensorboardX import SummaryWriter
def is_dist_avail_and_initialized():
if not dist.is_available():
return False
if not dist.is_initialized():
return False
return True
def get_world_size():
if not is_dist_avail_and_initialized():
return 1
return dist.get_world_size()
def get_rank():
if not is_dist_avail_and_initialized():
return 0
return dist.get_rank()
def is_main_process():
return get_rank() == 0
def save_on_master(*args, **kwargs):
if is_main_process():
torch.save(*args, **kwargs)
def setup_for_distributed(is_master):
"""
This function disables printing when not in master process
"""
import builtins as __builtin__
builtin_print = __builtin__.print
def print(*args, **kwargs):
force = kwargs.pop('force', False)
if is_master or force:
builtin_print(*args, **kwargs)
__builtin__.print = print
def init_distributed_mode(args):
if int(os.getenv('OMPI_COMM_WORLD_SIZE', '0')) > 0:
rank = int(os.environ['OMPI_COMM_WORLD_RANK'])
local_rank = int(os.environ['OMPI_COMM_WORLD_LOCAL_RANK'])
world_size = int(os.environ['OMPI_COMM_WORLD_SIZE'])
os.environ["LOCAL_RANK"] = os.environ['OMPI_COMM_WORLD_LOCAL_RANK']
os.environ["RANK"] = os.environ['OMPI_COMM_WORLD_RANK']
os.environ["WORLD_SIZE"] = os.environ['OMPI_COMM_WORLD_SIZE']
args.rank = int(os.environ["RANK"])
args.world_size = int(os.environ["WORLD_SIZE"])
args.gpu = int(os.environ["LOCAL_RANK"])
elif 'RANK' in os.environ and 'WORLD_SIZE' in os.environ:
args.rank = int(os.environ["RANK"])
args.world_size = int(os.environ['WORLD_SIZE'])
args.gpu = int(os.environ['LOCAL_RANK'])
else:
print('Not using distributed mode')
args.distributed = False
return
args.distributed = True
args.dist_backend = 'nccl'
args.dist_url = "env://"
print('| distributed init (rank {}): {}, gpu {}'.format(
args.rank, args.dist_url, args.gpu), flush=True)
def cosine_scheduler(base_value, final_value, epochs, niter_per_ep, warmup_epochs=0,
start_warmup_value=0, warmup_steps=-1):
warmup_schedule = np.array([])
warmup_iters = warmup_epochs * niter_per_ep
if warmup_steps > 0:
warmup_iters = warmup_steps
print("Set warmup steps = %d" % warmup_iters)
if warmup_epochs > 0:
warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters)
iters = np.arange(epochs * niter_per_ep - warmup_iters)
schedule = np.array(
[final_value + 0.5 * (base_value - final_value) * (1 + math.cos(math.pi * i / (len(iters)))) for i in iters])
schedule = np.concatenate((warmup_schedule, schedule))
assert len(schedule) == epochs * niter_per_ep
return schedule
def constant_scheduler(base_value, epochs, niter_per_ep, warmup_epochs=0,
start_warmup_value=1e-6, warmup_steps=-1):
warmup_schedule = np.array([])
warmup_iters = warmup_epochs * niter_per_ep
if warmup_steps > 0:
warmup_iters = warmup_steps
print("Set warmup steps = %d" % warmup_iters)
if warmup_iters > 0:
warmup_schedule = np.linspace(start_warmup_value, base_value, warmup_iters)
iters = epochs * niter_per_ep - warmup_iters
schedule = np.array([base_value] * iters)
schedule = np.concatenate((warmup_schedule, schedule))
assert len(schedule) == epochs * niter_per_ep
return schedule
def get_parameter_groups(model, weight_decay=1e-5, base_lr=1e-4, skip_list=(), get_num_layer=None, get_layer_scale=None, **kwargs):
parameter_group_names = {}
parameter_group_vars = {}
for name, param in model.named_parameters():
if not param.requires_grad:
continue # frozen weights
if len(kwargs.get('filter_name', [])) > 0:
flag = False
for filter_n in kwargs.get('filter_name', []):
if filter_n in name:
print(f"filter {name} because of the pattern {filter_n}")
flag = True
if flag:
continue
default_scale=1.
if param.ndim <= 1 or name.endswith(".bias") or name in skip_list: # param.ndim <= 1 len(param.shape) == 1
group_name = "no_decay"
this_weight_decay = 0.
else:
group_name = "decay"
this_weight_decay = weight_decay
if get_num_layer is not None:
layer_id = get_num_layer(name)
group_name = "layer_%d_%s" % (layer_id, group_name)
else:
layer_id = None
if group_name not in parameter_group_names:
if get_layer_scale is not None:
scale = get_layer_scale(layer_id)
else:
scale = default_scale
parameter_group_names[group_name] = {
"weight_decay": this_weight_decay,
"params": [],
"lr": base_lr,
"lr_scale": scale,
}
parameter_group_vars[group_name] = {
"weight_decay": this_weight_decay,
"params": [],
"lr": base_lr,
"lr_scale": scale,
}
parameter_group_vars[group_name]["params"].append(param)
parameter_group_names[group_name]["params"].append(name)
print("Param groups = %s" % json.dumps(parameter_group_names, indent=2))
return list(parameter_group_vars.values())
def create_optimizer(args, model, get_num_layer=None, get_layer_scale=None, filter_bias_and_bn=True, skip_list=None, **kwargs):
opt_lower = args.opt.lower()
weight_decay = args.weight_decay
skip = {}
if skip_list is not None:
skip = skip_list
elif hasattr(model, 'no_weight_decay'):
skip = model.no_weight_decay()
print(f"Skip weight decay name marked in model: {skip}")
parameters = get_parameter_groups(model, weight_decay, args.lr, skip, get_num_layer, get_layer_scale, **kwargs)
weight_decay = 0.
if 'fused' in opt_lower:
assert has_apex and torch.cuda.is_available(), 'APEX and CUDA required for fused optimizers'
opt_args = dict(lr=args.lr, weight_decay=weight_decay)
if hasattr(args, 'opt_eps') and args.opt_eps is not None:
opt_args['eps'] = args.opt_eps
if hasattr(args, 'opt_beta1') and args.opt_beta1 is not None:
opt_args['betas'] = (args.opt_beta1, args.opt_beta2)
print('Optimizer config:', opt_args)
opt_split = opt_lower.split('_')
opt_lower = opt_split[-1]
if opt_lower == 'sgd' or opt_lower == 'nesterov':
opt_args.pop('eps', None)
optimizer = optim.SGD(parameters, momentum=args.momentum, nesterov=True, **opt_args)
elif opt_lower == 'momentum':
opt_args.pop('eps', None)
optimizer = optim.SGD(parameters, momentum=args.momentum, nesterov=False, **opt_args)
elif opt_lower == 'adam':
optimizer = optim.Adam(parameters, **opt_args)
elif opt_lower == 'adamw':
optimizer = optim.AdamW(parameters, **opt_args)
elif opt_lower == 'adadelta':
optimizer = optim.Adadelta(parameters, **opt_args)
elif opt_lower == 'rmsprop':
optimizer = optim.RMSprop(parameters, alpha=0.9, momentum=args.momentum, **opt_args)
else:
assert False and "Invalid optimizer"
raise ValueError
return optimizer
class SmoothedValue(object):
"""Track a series of values and provide access to smoothed values over a
window or the global series average.
"""
def __init__(self, window_size=20, fmt=None):
if fmt is None:
fmt = "{median:.4f} ({global_avg:.4f})"
self.deque = deque(maxlen=window_size)
self.total = 0.0
self.count = 0
self.fmt = fmt
def update(self, value, n=1):
self.deque.append(value)
self.count += n
self.total += value * n
def synchronize_between_processes(self):
"""
Warning: does not synchronize the deque!
"""
if not is_dist_avail_and_initialized():
return
t = torch.tensor([self.count, self.total], dtype=torch.float64, device='cuda')
dist.barrier()
dist.all_reduce(t)
t = t.tolist()
self.count = int(t[0])
self.total = t[1]
@property
def median(self):
d = torch.tensor(list(self.deque))
return d.median().item()
@property
def avg(self):
d = torch.tensor(list(self.deque), dtype=torch.float32)
return d.mean().item()
@property
def global_avg(self):
return self.total / self.count
@property
def max(self):
return max(self.deque)
@property
def value(self):
return self.deque[-1]
def __str__(self):
return self.fmt.format(
median=self.median,
avg=self.avg,
global_avg=self.global_avg,
max=self.max,
value=self.value)
class MetricLogger(object):
def __init__(self, delimiter="\t"):
self.meters = defaultdict(SmoothedValue)
self.delimiter = delimiter
def update(self, **kwargs):
for k, v in kwargs.items():
if v is None:
continue
if isinstance(v, torch.Tensor):
v = v.item()
assert isinstance(v, (float, int))
self.meters[k].update(v)
def __getattr__(self, attr):
if attr in self.meters:
return self.meters[attr]
if attr in self.__dict__:
return self.__dict__[attr]
raise AttributeError("'{}' object has no attribute '{}'".format(
type(self).__name__, attr))
def __str__(self):
loss_str = []
for name, meter in self.meters.items():
loss_str.append(
"{}: {}".format(name, str(meter))
)
return self.delimiter.join(loss_str)
def synchronize_between_processes(self):
for meter in self.meters.values():
meter.synchronize_between_processes()
def add_meter(self, name, meter):
self.meters[name] = meter
def log_every(self, iterable, print_freq, header=None):
i = 0
if not header:
header = ''
start_time = time.time()
end = time.time()
iter_time = SmoothedValue(fmt='{avg:.4f}')
data_time = SmoothedValue(fmt='{avg:.4f}')
space_fmt = ':' + str(len(str(len(iterable)))) + 'd'
log_msg = [
header,
'[{0' + space_fmt + '}/{1}]',
'eta: {eta}',
'{meters}',
'time: {time}',
'data: {data}'
]
if torch.cuda.is_available():
log_msg.append('max mem: {memory:.0f}')
log_msg = self.delimiter.join(log_msg)
MB = 1024.0 * 1024.0
for obj in iterable:
data_time.update(time.time() - end)
yield obj
iter_time.update(time.time() - end)
if i % print_freq == 0 or i == len(iterable) - 1:
eta_seconds = iter_time.global_avg * (len(iterable) - i)
eta_string = str(datetime.timedelta(seconds=int(eta_seconds)))
if torch.cuda.is_available():
print(log_msg.format(
i, len(iterable), eta=eta_string,
meters=str(self),
time=str(iter_time), data=str(data_time),
memory=torch.cuda.max_memory_allocated() / MB))
else:
print(log_msg.format(
i, len(iterable), eta=eta_string,
meters=str(self),
time=str(iter_time), data=str(data_time)))
i += 1
end = time.time()
total_time = time.time() - start_time
total_time_str = str(datetime.timedelta(seconds=int(total_time)))
print('{} Total time: {} ({:.4f} s / it)'.format(
header, total_time_str, total_time / len(iterable)))
+1 -1
View File
@@ -117,7 +117,7 @@ class CausalVideoVAE(ModelMixin, ConfigMixin):
):
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
self.encoder = CausalVaeEncoder(