Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
426e365304 | ||
|
|
3b572e270b | ||
|
|
50f5d1fb37 | ||
|
|
723537bb75 | ||
|
|
acb19bcd29 | ||
|
|
1d9fe8c695 | ||
|
|
3aeea32237 | ||
|
|
bc8d4360fe | ||
|
|
ecd00cf903 | ||
|
|
eb5668b88b | ||
|
|
7387ea5de5 | ||
|
|
3e14b0d77e |
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+511
-527
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
}
|
||||
+227
-189
@@ -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
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,3 +1,2 @@
|
||||
from .modeling_pyramid_flux import PyramidFluxTransformer
|
||||
from .modeling_text_encoder import FluxTextEncoderWithMask
|
||||
from .modeling_flux_block import FluxSingleTransformerBlock, FluxTransformerBlock
|
||||
@@ -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)
|
||||
|
||||
@@ -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,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:]
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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)))
|
||||
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user