diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml new file mode 100644 index 0000000..828f300 --- /dev/null +++ b/.github/workflows/publish.yml @@ -0,0 +1,21 @@ +name: Publish to Comfy registry +on: + workflow_dispatch: + push: + branches: + - main + paths: + - "pyproject.toml" + +jobs: + publish-node: + name: Publish Custom Node to registry + runs-on: ubuntu-latest + steps: + - name: Check out code + uses: actions/checkout@v4 + - name: Publish Custom Node + uses: Comfy-Org/publish-node-action@main + with: + ## Add your own personal access token to your Github Repository secrets and reference it here. + personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }} \ No newline at end of file diff --git a/README.md b/README.md index 1a977bd..d0d5c48 100644 --- a/README.md +++ b/README.md @@ -1,4 +1,10 @@ ## DynamiCrafter wrapper nodes for ComfyUI + +## Update2: Refactor + +Changed lots of things to better integrate this to ComfyUI, you can (and have to) use clip_vision and clip models, but memory usage is much better and I was able to do 512x320 under 10GB VRAM. +New example workflows are included, all old workflows will have to be updated. + ## Update: ToonCrafter Initial ToonCrafter support with it's own node. diff --git a/examples/DynamiCrafter-CIL-testing_example_01.json b/examples/DynamiCrafter-CIL-testing_example_01.json new file mode 100644 index 0000000..9b8a7d8 --- /dev/null +++ b/examples/DynamiCrafter-CIL-testing_example_01.json @@ -0,0 +1,1240 @@ +{ + "last_node_id": 74, + "last_link_id": 185, + "nodes": [ + { + "id": 66, + "type": "AddLabel", + "pos": [ + 2050, + 160 + ], + "size": { + "0": 315, + "1": 274 + }, + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 163 + }, + { + "name": "caption", + "type": "STRING", + "link": null, + "widget": { + "name": "caption" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 164 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "AddLabel" + }, + "widgets_values": [ + 10, + 2, + 48, + 32, + "white", + "black", + "FreeMono.ttf", + "analytic_init_noise", + "up", + "" + ] + }, + { + "id": 1, + "type": "LoadImage", + "pos": [ + 490, + 200 + ], + "size": { + "0": 315, + "1": 314 + }, + "flags": {}, + "order": 0, + "mode": 0, + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 2 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "MASK", + "type": "MASK", + "links": [], + "shape": 3, + "slot_index": 1 + } + ], + "properties": { + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "Mona-Lisa-oil-wood-panel-Leonardo-da.webp", + "image" + ] + }, + { + "id": 52, + "type": "DownloadAndLoadDynamiCrafterModel", + "pos": [ + 534, + -178 + ], + "size": { + "0": 433.2352294921875, + "1": 106 + }, + "flags": {}, + "order": 1, + "mode": 0, + "outputs": [ + { + "name": "DynCraft_model", + "type": "DCMODEL", + "links": [ + 138, + 152, + 155 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadDynamiCrafterModel" + }, + "widgets_values": [ + "dynamicrafter-CIL-512-no-watermark-pruned-fp16.safetensors", + "auto", + false + ] + }, + { + "id": 64, + "type": "ImageConcanate", + "pos": [ + 2040, + -40 + ], + "size": { + "0": 315, + "1": 102 + }, + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "image1", + "type": "IMAGE", + "link": 162 + }, + { + "name": "image2", + "type": "IMAGE", + "link": 164 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 165 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "ImageConcanate" + }, + "widgets_values": [ + "right", + false + ] + }, + { + "id": 60, + "type": "DownloadAndLoadCLIPModel", + "pos": [ + 574, + -305 + ], + "size": { + "0": 371.02264404296875, + "1": 64.01405334472656 + }, + "flags": {}, + "order": 2, + "mode": 0, + "outputs": [ + { + "name": "clip", + "type": "CLIP", + "links": [ + 147, + 148 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadCLIPModel" + }, + "widgets_values": [ + "stable-diffusion-2-1-clip-fp16.safetensors" + ] + }, + { + "id": 50, + "type": "CLIPTextEncode", + "pos": [ + 1025, + -65 + ], + "size": { + "0": 400, + "1": 200 + }, + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 148 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 141, + 158 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "" + ] + }, + { + "id": 29, + "type": "VHS_VideoCombine", + "pos": [ + 2410, + -390 + ], + "size": [ + 1550.3211669921875, + 853.9591693878174 + ], + "flags": {}, + "order": 17, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 165 + }, + { + "name": "audio", + "type": "VHS_AUDIO", + "link": null + }, + { + "name": "meta_batch", + "type": "VHS_BatchManager", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 8, + "loop_count": 0, + "filename_prefix": "AnimateDiff", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "AnimateDiff_00005.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4" + } + } + } + }, + { + "id": 62, + "type": "DynamiCrafterLoadInitNoise", + "pos": [ + 495, + 1 + ], + "size": { + "0": 315, + "1": 122 + }, + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "DCMODEL", + "link": 152 + } + ], + "outputs": [ + { + "name": "init_noise", + "type": "DCNOISE", + "links": [ + 154 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": [ + 170 + ], + "shape": 3, + "slot_index": 1 + }, + { + "name": "height", + "type": "INT", + "links": [ + 171 + ], + "shape": 3, + "slot_index": 2 + } + ], + "properties": { + "Node name for S&R": "DynamiCrafterLoadInitNoise" + }, + "widgets_values": [ + 940, + true + ] + }, + { + "id": 5, + "type": "ImageResizeKJ", + "pos": [ + 973, + 239 + ], + "size": { + "0": 315, + "1": 242 + }, + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 2 + }, + { + "name": "get_image_size", + "type": "IMAGE", + "link": null + }, + { + "name": "width_input", + "type": "INT", + "link": 170, + "widget": { + "name": "width_input" + } + }, + { + "name": "height_input", + "type": "INT", + "link": 171, + "widget": { + "name": "height_input" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 172, + 173 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "width", + "type": "INT", + "links": null, + "shape": 3 + }, + { + "name": "height", + "type": "INT", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "ImageResizeKJ" + }, + "widgets_values": [ + 512, + 320, + "lanczos", + false, + 64, + 0, + 0 + ] + }, + { + "id": 49, + "type": "CLIPTextEncode", + "pos": [ + 1028, + -310 + ], + "size": { + "0": 400, + "1": 200 + }, + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 147, + "slot_index": 0 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "links": [ + 140, + 157 + ], + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "nodding" + ] + }, + { + "id": 59, + "type": "DownloadAndLoadCLIPVisionModel", + "pos": [ + 622, + -439 + ], + "size": { + "0": 315, + "1": 58 + }, + "flags": {}, + "order": 3, + "mode": 0, + "outputs": [ + { + "name": "clip_vision", + "type": "CLIP_VISION", + "links": [ + 146, + 156 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "DownloadAndLoadCLIPVisionModel" + }, + "widgets_values": [ + "CLIP-ViT-H-fp16.safetensors" + ] + }, + { + "id": 65, + "type": "AddLabel", + "pos": [ + 2060, + -390 + ], + "size": { + "0": 315, + "1": 274 + }, + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 161 + }, + { + "name": "caption", + "type": "STRING", + "link": null, + "widget": { + "name": "caption" + } + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 162 + ], + "shape": 3, + "slot_index": 0 + } + ], + "properties": { + "Node name for S&R": "AddLabel" + }, + "widgets_values": [ + 10, + 2, + 48, + 32, + "white", + "black", + "FreeMono.ttf", + "baseline", + "up", + "" + ] + }, + { + "id": 71, + "type": "PrimitiveNode", + "pos": [ + 1332, + -736 + ], + "size": [ + 245.37998624877855, + 82 + ], + "flags": {}, + "order": 4, + "mode": 0, + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 175, + 176 + ], + "widget": { + "name": "seed" + }, + "slot_index": 0 + } + ], + "title": "seed", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 619731667089950, + "fixed" + ] + }, + { + "id": 72, + "type": "PrimitiveNode", + "pos": [ + 1327, + -601 + ], + "size": [ + 252.17814718627847, + 82 + ], + "flags": {}, + "order": 5, + "mode": 0, + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 178, + 179 + ], + "widget": { + "name": "steps" + } + } + ], + "title": "steps", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 26, + "fixed" + ] + }, + { + "id": 73, + "type": "PrimitiveNode", + "pos": [ + 1328, + -466 + ], + "size": { + "0": 210, + "1": 82 + }, + "flags": {}, + "order": 6, + "mode": 0, + "outputs": [ + { + "name": "FLOAT", + "type": "FLOAT", + "links": [ + 181, + 182 + ], + "widget": { + "name": "cfg" + }, + "slot_index": 0 + } + ], + "title": "cfg", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 7, + "fixed" + ] + }, + { + "id": 63, + "type": "DynamiCrafterI2V", + "pos": [ + 1670, + -390 + ], + "size": [ + 315, + 462 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "DCMODEL", + "link": 155 + }, + { + "name": "clip_vision", + "type": "CLIP_VISION", + "link": 156 + }, + { + "name": "positive", + "type": "CONDITIONING", + "link": 157 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 158 + }, + { + "name": "image", + "type": "IMAGE", + "link": 173 + }, + { + "name": "image2", + "type": "IMAGE", + "link": null + }, + { + "name": "mask", + "type": "MASK", + "link": null + }, + { + "name": "init_noise", + "type": "DCNOISE", + "link": null, + "slot_index": 7 + }, + { + "name": "seed", + "type": "INT", + "link": 176, + "widget": { + "name": "seed" + } + }, + { + "name": "steps", + "type": "INT", + "link": 178, + "widget": { + "name": "steps" + }, + "slot_index": 9 + }, + { + "name": "cfg", + "type": "FLOAT", + "link": 182, + "widget": { + "name": "cfg" + } + }, + { + "name": "fs", + "type": "INT", + "link": 185, + "widget": { + "name": "fs" + } + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 161 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "last_image", + "type": "IMAGE", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "DynamiCrafterI2V" + }, + "widgets_values": [ + 26, + 7, + 1, + 16, + 619731667089950, + "fixed", + 24, + true, + "auto", + 16, + 4, + 0 + ] + }, + { + "id": 58, + "type": "DynamiCrafterI2V", + "pos": [ + 1670, + 150 + ], + "size": [ + 315, + 462 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "DCMODEL", + "link": 138 + }, + { + "name": "clip_vision", + "type": "CLIP_VISION", + "link": 146 + }, + { + "name": "positive", + "type": "CONDITIONING", + "link": 140 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 141 + }, + { + "name": "image", + "type": "IMAGE", + "link": 172 + }, + { + "name": "image2", + "type": "IMAGE", + "link": null + }, + { + "name": "mask", + "type": "MASK", + "link": null + }, + { + "name": "init_noise", + "type": "DCNOISE", + "link": 154, + "slot_index": 7 + }, + { + "name": "seed", + "type": "INT", + "link": 175, + "widget": { + "name": "seed" + }, + "slot_index": 8 + }, + { + "name": "steps", + "type": "INT", + "link": 179, + "widget": { + "name": "steps" + }, + "slot_index": 9 + }, + { + "name": "cfg", + "type": "FLOAT", + "link": 181, + "widget": { + "name": "cfg" + }, + "slot_index": 10 + }, + { + "name": "fs", + "type": "INT", + "link": 184, + "widget": { + "name": "fs" + }, + "slot_index": 11 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 163 + ], + "shape": 3, + "slot_index": 0 + }, + { + "name": "last_image", + "type": "IMAGE", + "links": null, + "shape": 3 + } + ], + "properties": { + "Node name for S&R": "DynamiCrafterI2V" + }, + "widgets_values": [ + 26, + 7, + 1, + 16, + 619731667089950, + "fixed", + 24, + true, + "auto", + 16, + 4, + 0 + ] + }, + { + "id": 74, + "type": "PrimitiveNode", + "pos": [ + 1614, + -727 + ], + "size": { + "0": 210, + "1": 82 + }, + "flags": {}, + "order": 7, + "mode": 0, + "outputs": [ + { + "name": "INT", + "type": "INT", + "links": [ + 184, + 185 + ], + "widget": { + "name": "fs" + }, + "slot_index": 0 + } + ], + "title": "fs", + "properties": { + "Run widget replace on values": false + }, + "widgets_values": [ + 24, + "fixed" + ] + } + ], + "links": [ + [ + 2, + 1, + 0, + 5, + 0, + "IMAGE" + ], + [ + 138, + 52, + 0, + 58, + 0, + "DCMODEL" + ], + [ + 140, + 49, + 0, + 58, + 2, + "CONDITIONING" + ], + [ + 141, + 50, + 0, + 58, + 3, + "CONDITIONING" + ], + [ + 146, + 59, + 0, + 58, + 1, + "CLIP_VISION" + ], + [ + 147, + 60, + 0, + 49, + 0, + "CLIP" + ], + [ + 148, + 60, + 0, + 50, + 0, + "CLIP" + ], + [ + 152, + 52, + 0, + 62, + 0, + "DCMODEL" + ], + [ + 154, + 62, + 0, + 58, + 7, + "DCNOISE" + ], + [ + 155, + 52, + 0, + 63, + 0, + "DCMODEL" + ], + [ + 156, + 59, + 0, + 63, + 1, + "CLIP_VISION" + ], + [ + 157, + 49, + 0, + 63, + 2, + "CONDITIONING" + ], + [ + 158, + 50, + 0, + 63, + 3, + "CONDITIONING" + ], + [ + 161, + 63, + 0, + 65, + 0, + "IMAGE" + ], + [ + 162, + 65, + 0, + 64, + 0, + "IMAGE" + ], + [ + 163, + 58, + 0, + 66, + 0, + "IMAGE" + ], + [ + 164, + 66, + 0, + 64, + 1, + "IMAGE" + ], + [ + 165, + 64, + 0, + 29, + 0, + "IMAGE" + ], + [ + 170, + 62, + 1, + 5, + 2, + "INT" + ], + [ + 171, + 62, + 2, + 5, + 3, + "INT" + ], + [ + 172, + 5, + 0, + 58, + 4, + "IMAGE" + ], + [ + 173, + 5, + 0, + 63, + 4, + "IMAGE" + ], + [ + 175, + 71, + 0, + 58, + 8, + "INT" + ], + [ + 176, + 71, + 0, + 63, + 8, + "INT" + ], + [ + 178, + 72, + 0, + 63, + 9, + "INT" + ], + [ + 179, + 72, + 0, + 58, + 9, + "INT" + ], + [ + 181, + 73, + 0, + 58, + 10, + "FLOAT" + ], + [ + 182, + 73, + 0, + 63, + 10, + "FLOAT" + ], + [ + 184, + 74, + 0, + 58, + 11, + "INT" + ], + [ + 185, + 74, + 0, + 63, + 11, + "INT" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 0.6830134553650712, + "offset": { + "0": -357.3632507324219, + "1": 839.998046875 + } + } + }, + "version": 0.4 +} \ No newline at end of file diff --git a/examples/dynamicrafter_i2v_example_01.json b/examples/dynamicrafter_i2v_example_01.json index 51e9bcf..a495cdd 100644 --- a/examples/dynamicrafter_i2v_example_01.json +++ b/examples/dynamicrafter_i2v_example_01.json @@ -431,7 +431,7 @@ "Node name for S&R": "DownloadAndLoadDynamiCrafterModel" }, "widgets_values": [ - "dynamicrafter_1024_v1_bf16.safetensors", + "dynamicrafter_1024_fp16_pruned.safetensors", "auto", true ] diff --git a/examples/tooncrafter_example_01.json b/examples/tooncrafter_example_01.json index 67910e6..321f563 100644 --- a/examples/tooncrafter_example_01.json +++ b/examples/tooncrafter_example_01.json @@ -617,7 +617,7 @@ "Node name for S&R": "DownloadAndLoadDynamiCrafterModel" }, "widgets_values": [ - "tooncrafter_512_interp-fp16.safetensors", + "tooncrafter_512_interp-pruned-fp16.safetensors", "auto", false ] diff --git a/examples/tooncrafter_example_low_vram_01.json b/examples/tooncrafter_example_low_vram_01.json index ac04851..25b9481 100644 --- a/examples/tooncrafter_example_low_vram_01.json +++ b/examples/tooncrafter_example_low_vram_01.json @@ -549,7 +549,7 @@ "Node name for S&R": "DownloadAndLoadDynamiCrafterModel" }, "widgets_values": [ - "tooncrafter_512_interp-fp16.safetensors", + "tooncrafter_512_interp-pruned-fp16.safetensors", "auto", false ] diff --git a/init_noises/initial_noise_1024.safetensors b/init_noises/initial_noise_1024.safetensors new file mode 100644 index 0000000..0f3ebf4 Binary files /dev/null and b/init_noises/initial_noise_1024.safetensors differ diff --git a/init_noises/initial_noise_512.safetensors b/init_noises/initial_noise_512.safetensors new file mode 100644 index 0000000..4bf49cf Binary files /dev/null and b/init_noises/initial_noise_512.safetensors differ diff --git a/lvdm/basics.py b/lvdm/basics.py index 527e29f..2cd39be 100644 --- a/lvdm/basics.py +++ b/lvdm/basics.py @@ -9,7 +9,8 @@ import torch.nn as nn from ..utils.utils import instantiate_from_config - +import comfy.ops +ops = comfy.ops.manual_cast def disabled_train(self, mode=True): """Overwrite model.train with this function to make sure train/eval mode @@ -38,11 +39,11 @@ def conv_nd(dims, *args, **kwargs): Create a 1D, 2D, or 3D convolution module. """ if dims == 1: - return nn.Conv1d(*args, **kwargs) + return ops.Conv1d(*args, **kwargs) elif dims == 2: - return nn.Conv2d(*args, **kwargs) + return ops.Conv2d(*args, **kwargs) elif dims == 3: - return nn.Conv3d(*args, **kwargs) + return ops.Conv3d(*args, **kwargs) raise ValueError(f"unsupported dimensions: {dims}") @@ -50,7 +51,7 @@ def linear(*args, **kwargs): """ Create a linear module. """ - return nn.Linear(*args, **kwargs) + return ops.Linear(*args, **kwargs) def avg_pool_nd(dims, *args, **kwargs): diff --git a/lvdm/models/autoencoder.py b/lvdm/models/autoencoder.py index baa736e..807f800 100644 --- a/lvdm/models/autoencoder.py +++ b/lvdm/models/autoencoder.py @@ -8,7 +8,8 @@ import pytorch_lightning as pl from ...lvdm.modules.networks.ae_modules import Encoder, Decoder from ...lvdm.distributions import DiagonalGaussianDistribution from ...utils.utils import instantiate_from_config - +import comfy.ops +ops = comfy.ops.manual_cast TIMESTEPS=16 class AutoencoderKL(pl.LightningModule): def __init__(self, @@ -34,8 +35,8 @@ class AutoencoderKL(pl.LightningModule): self.decoder = Decoder(**ddconfig) self.loss = instantiate_from_config(lossconfig) assert ddconfig["double_z"] - self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) - self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) + self.quant_conv = ops.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) + self.post_quant_conv = ops.Conv2d(embed_dim, ddconfig["z_channels"], 1) self.embed_dim = embed_dim self.input_dim = input_dim self.test = test diff --git a/lvdm/models/autoencoder_dualref.py b/lvdm/models/autoencoder_dualref.py index 55e8db1..8dde8ea 100644 --- a/lvdm/models/autoencoder_dualref.py +++ b/lvdm/models/autoencoder_dualref.py @@ -9,17 +9,23 @@ import torch.nn as nn from packaging import version logpy = logging.getLogger(__name__) -try: - import xformers - import xformers.ops +import comfy.model_management +if comfy.model_management.XFORMERS_IS_AVAILABLE: + try: + import xformers + import xformers.ops - XFORMERS_IS_AVAILABLE = True -except: + XFORMERS_IS_AVAILABLE = True + except: + XFORMERS_IS_AVAILABLE = False + logpy.warning("no module 'xformers'. Processing without...") +else: XFORMERS_IS_AVAILABLE = False - logpy.warning("no module 'xformers'. Processing without...") -from ...lvdm.modules.attention_svd import LinearAttention, MemoryEfficientCrossAttention +from ...lvdm.modules.attention_svd import LinearAttention, MemoryEfficientCrossAttention, CrossAttention +import comfy.ops +ops = comfy.ops.manual_cast def nonlinearity(x): # swish @@ -27,7 +33,7 @@ def nonlinearity(x): def Normalize(in_channels, num_groups=32): - return torch.nn.GroupNorm( + return ops.GroupNorm( num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True ) @@ -49,23 +55,23 @@ class ResnetBlock(nn.Module): self.use_conv_shortcut = conv_shortcut self.norm1 = Normalize(in_channels) - self.conv1 = torch.nn.Conv2d( + self.conv1 = ops.Conv2d( in_channels, out_channels, kernel_size=3, stride=1, padding=1 ) if temb_channels > 0: - self.temb_proj = torch.nn.Linear(temb_channels, out_channels) + self.temb_proj = ops.Linear(temb_channels, out_channels) self.norm2 = Normalize(out_channels) self.dropout = torch.nn.Dropout(dropout) - self.conv2 = torch.nn.Conv2d( + self.conv2 = ops.Conv2d( out_channels, out_channels, kernel_size=3, stride=1, padding=1 ) if self.in_channels != self.out_channels: if self.use_conv_shortcut: - self.conv_shortcut = torch.nn.Conv2d( + self.conv_shortcut = ops.Conv2d( in_channels, out_channels, kernel_size=3, stride=1, padding=1 ) else: - self.nin_shortcut = torch.nn.Conv2d( + self.nin_shortcut = ops.Conv2d( in_channels, out_channels, kernel_size=1, stride=1, padding=0 ) @@ -105,16 +111,16 @@ class AttnBlock(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d( + self.q = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.k = torch.nn.Conv2d( + self.k = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.v = torch.nn.Conv2d( + self.v = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.proj_out = torch.nn.Conv2d( + self.proj_out = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) @@ -155,16 +161,16 @@ class MemoryEfficientAttnBlock(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d( + self.q = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.k = torch.nn.Conv2d( + self.k = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.v = torch.nn.Conv2d( + self.v = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.proj_out = torch.nn.Conv2d( + self.proj_out = ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) self.attention_op: Optional[Any] = None @@ -206,6 +212,14 @@ class MemoryEfficientAttnBlock(nn.Module): return x + h_ +class CrossAttentionWrapper(CrossAttention): + def forward(self, x, context=None, mask=None, **unused_kwargs): + b, c, h, w = x.shape + x = rearrange(x, "b c h w -> b 1 (h w) c").contiguous() + out = super().forward(x, context=context, mask=mask) + out = rearrange(out, "b 1 (h w) c -> b c h w", h=h, w=w, c=c, b=b) + return x + out + class MemoryEfficientCrossAttentionWrapper(MemoryEfficientCrossAttention): def forward(self, x, context=None, mask=None, **unused_kwargs): b, c, h, w = x.shape @@ -220,9 +234,11 @@ def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None): assert attn_type in [ "vanilla", "vanilla-xformers", + "cross-attn", "memory-efficient-cross-attn", "linear", "none", + "cross-attn-fusion", "memory-efficient-cross-attn-fusion", ], f"attn_type {attn_type} unknown" if ( @@ -243,9 +259,15 @@ def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None): f"building MemoryEfficientAttnBlock with {in_channels} in_channels..." ) return MemoryEfficientAttnBlock(in_channels) + elif attn_type == "cross-attn": + attn_kwargs["query_dim"] = in_channels + return CrossAttentionWrapper(**attn_kwargs) elif attn_type == "memory-efficient-cross-attn": attn_kwargs["query_dim"] = in_channels return MemoryEfficientCrossAttentionWrapper(**attn_kwargs) + elif attn_type == "cross-attn-fusion": + attn_kwargs["query_dim"] = in_channels + return CrossAttentionWrapperFusion(**attn_kwargs) elif attn_type == "memory-efficient-cross-attn-fusion": attn_kwargs["query_dim"] = in_channels return MemoryEfficientCrossAttentionWrapperFusion(**attn_kwargs) @@ -254,6 +276,76 @@ def make_attn(in_channels, attn_type="vanilla", attn_kwargs=None): else: return LinAttnBlock(in_channels) + +class CrossAttentionWrapperFusion(CrossAttention): + def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0, **kwargs): + super().__init__(query_dim, context_dim, heads, dim_head, dropout, **kwargs) + self.dim_head = dim_head + self.norm = Normalize(query_dim) + nn.init.zeros_(self.to_out[0].weight) + nn.init.zeros_(self.to_out[0].bias) + + def forward(self, x, context=None, mask=None): + if self.training: + return checkpoint(self._forward, x, context, mask, use_reentrant=False) + else: + return self._forward(x, context, mask) + + def _forward( + self, + x, + context=None, + mask=None, + ): + bt, c, h, w = x.shape + h_ = self.norm(x) + h_ = rearrange(h_, "b c h w -> b (h w) c") + q = self.to_q(h_) + + b, c, l, h, w = context.shape + context = rearrange(context, "b c l h w -> (b l) (h w) c") + k = self.to_k(context) + v = self.to_v(context) + k = rearrange(k, "(b l) d c -> b l d c", l=l) + k = torch.cat([k[:, [0] * (bt // b)], k[:, [1] * (bt // b)]], dim=2) + k = rearrange(k, "b l d c -> (b l) d c") + + v = rearrange(v, "(b l) d c -> b l d c", l=l) + v = torch.cat([v[:, [0] * (bt // b)], v[:, [1] * (bt // b)]], dim=2) + v = rearrange(v, "b l d c -> (b l) d c") + + b, _, _ = q.shape # actually bt + q, k, v = map( + lambda t: t.unsqueeze(3) + .reshape(b, t.shape[1], self.heads, self.dim_head) + .permute(0, 2, 1, 3) + .reshape(b * self.heads, t.shape[1], self.dim_head) + .contiguous(), + (q, k, v), + ) + sdpa = torch.nn.functional.scaled_dot_product_attention + + def slow_sdpa(q, k, v): + out_list = [] + step = 10 + for i in range(0, q.shape[0], step): + out_i = sdpa(q[i:i + step], k[i:i + step], v[i:i + step]) + out_list.append(out_i) + return torch.cat(out_list, dim=0) + + out = slow_sdpa(q, k, v) + + out = ( + out.unsqueeze(0) + .reshape(b, self.heads, out.shape[1], self.dim_head) + .permute(0, 2, 1, 3) + .reshape(b, out.shape[1], self.heads * self.dim_head) + ) + out = self.to_out(out) + out = rearrange(out, "bt (h w) c -> bt c h w", h=h, w=w, c=c) + return x + out + + class MemoryEfficientCrossAttentionWrapperFusion(MemoryEfficientCrossAttention): # print('x.shape: ',x.shape, 'context.shape: ',context.shape) ##torch.Size([8, 128, 256, 256]) torch.Size([1, 128, 2, 256, 256]) def __init__(self, query_dim, context_dim=None, heads=8, dim_head=64, dropout=0, **kwargs): @@ -344,7 +436,7 @@ class MemoryEfficientCrossAttentionWrapperFusion(MemoryEfficientCrossAttention): class Combiner(nn.Module): def __init__(self, ch) -> None: super().__init__() - self.conv = nn.Conv2d(ch,ch,1,padding=0) + self.conv = ops.Conv2d(ch,ch,1,padding=0) nn.init.zeros_(self.conv.weight) nn.init.zeros_(self.conv.bias) @@ -417,7 +509,7 @@ class Decoder(nn.Module): make_resblock_cls = self._make_resblock() make_conv_cls = self._make_conv() # z to block_in - self.conv_in = torch.nn.Conv2d( + self.conv_in = ops.Conv2d( z_channels, block_in, kernel_size=3, stride=1, padding=1 ) @@ -465,7 +557,8 @@ class Decoder(nn.Module): self.up.insert(0, up) # prepend to get consistent order if i_level in self.attn_level: - self.attn_refinement.insert(0, make_attn_cls(block_in, attn_type='memory-efficient-cross-attn-fusion', attn_kwargs={})) + _attn_type = 'memory-efficient-cross-attn-fusion' if XFORMERS_IS_AVAILABLE else 'cross-attn-fusion' + self.attn_refinement.insert(0, make_attn_cls(block_in, attn_type=_attn_type, attn_kwargs={})) else: self.attn_refinement.insert(0, Combiner(block_in)) # end @@ -482,7 +575,7 @@ class Decoder(nn.Module): return ResnetBlock def _make_conv(self) -> Callable: - return torch.nn.Conv2d + return ops.Conv2d def get_last_layer(self, **kwargs): return self.conv_out.weight @@ -737,7 +830,7 @@ class VideoTransformerBlock(nn.Module): self.is_res = inner_dim == dim if self.ff_in: - self.norm_in = nn.LayerNorm(dim) + self.norm_in = ops.LayerNorm(dim) self.ff_in = FeedForward( dim, dim_out=inner_dim, dropout=dropout, glu=gated_ff ) @@ -765,7 +858,7 @@ class VideoTransformerBlock(nn.Module): else: self.attn2 = None else: - self.norm2 = nn.LayerNorm(inner_dim) + self.norm2 = ops.LayerNorm(inner_dim) if switch_temporal_ca_to_sa: self.attn2 = attn_cls( query_dim=inner_dim, heads=n_heads, dim_head=d_head, dropout=dropout @@ -779,8 +872,8 @@ class VideoTransformerBlock(nn.Module): dropout=dropout, ) # is self-attn if context is none - self.norm1 = nn.LayerNorm(inner_dim) - self.norm3 = nn.LayerNorm(inner_dim) + self.norm1 = ops.LayerNorm(inner_dim) + self.norm3 = ops.LayerNorm(inner_dim) self.switch_temporal_ca_to_sa = switch_temporal_ca_to_sa self.checkpoint = checkpoint @@ -912,7 +1005,7 @@ class VideoResBlock(ResnetBlock): return x -class AE3DConv(torch.nn.Conv2d): +class AE3DConv(ops.Conv2d): def __init__(self, in_channels, out_channels, video_kernel_size=3, *args, **kwargs): super().__init__(in_channels, out_channels, *args, **kwargs) if isinstance(video_kernel_size, Iterable): @@ -920,7 +1013,7 @@ class AE3DConv(torch.nn.Conv2d): else: padding = int(video_kernel_size // 2) - self.time_mix_conv = torch.nn.Conv3d( + self.time_mix_conv = ops.Conv3d( in_channels=out_channels, out_channels=out_channels, kernel_size=video_kernel_size, @@ -953,9 +1046,9 @@ class VideoBlock(AttnBlock): time_embed_dim = self.in_channels * 4 self.video_time_embed = torch.nn.Sequential( - torch.nn.Linear(self.in_channels, time_embed_dim), + ops.Linear(self.in_channels, time_embed_dim), torch.nn.SiLU(), - torch.nn.Linear(time_embed_dim, self.in_channels), + ops.Linear(time_embed_dim, self.in_channels), ) self.merge_strategy = merge_strategy @@ -1023,9 +1116,9 @@ class MemoryEfficientVideoBlock(MemoryEfficientAttnBlock): time_embed_dim = self.in_channels * 4 self.video_time_embed = torch.nn.Sequential( - torch.nn.Linear(self.in_channels, time_embed_dim), + ops.Linear(self.in_channels, time_embed_dim), torch.nn.SiLU(), - torch.nn.Linear(time_embed_dim, self.in_channels), + ops.Linear(time_embed_dim, self.in_channels), ) self.merge_strategy = merge_strategy @@ -1114,7 +1207,7 @@ def make_time_attn( return NotImplementedError() -class Conv2DWrapper(torch.nn.Conv2d): +class Conv2DWrapper(ops.Conv2d): def forward(self, input: torch.Tensor, **kwargs) -> torch.Tensor: return super().forward(input) diff --git a/lvdm/models/autoencoder_old.py b/lvdm/models/autoencoder_old.py index d49e04e..5f21013 100644 --- a/lvdm/models/autoencoder_old.py +++ b/lvdm/models/autoencoder_old.py @@ -9,6 +9,9 @@ from ...lvdm.modules.networks.ae_modules import Encoder, Decoder from ...lvdm.distributions import DiagonalGaussianDistribution from ...utils.utils import instantiate_from_config +import comfy.ops +ops = comfy.ops.manual_cast + TIMESTEPS=16 class AutoencoderKL(pl.LightningModule): def __init__(self, @@ -31,8 +34,8 @@ class AutoencoderKL(pl.LightningModule): self.decoder = Decoder(**ddconfig) self.loss = instantiate_from_config(lossconfig) assert ddconfig["double_z"] - self.quant_conv = torch.nn.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) - self.post_quant_conv = torch.nn.Conv2d(embed_dim, ddconfig["z_channels"], 1) + self.quant_conv = ops.Conv2d(2*ddconfig["z_channels"], 2*embed_dim, 1) + self.post_quant_conv = ops.Conv2d(embed_dim, ddconfig["z_channels"], 1) self.embed_dim = embed_dim self.input_dim = input_dim self.test = test diff --git a/lvdm/models/ddpm3d.py b/lvdm/models/ddpm3d.py index 77e2d25..668b87d 100644 --- a/lvdm/models/ddpm3d.py +++ b/lvdm/models/ddpm3d.py @@ -33,6 +33,8 @@ from ...lvdm.models.autoencoder_dualref import VideoDecoder __conditioning_keys__ = {'concat': 'c_concat', 'crossattn': 'c_crossattn', 'adm': 'y'} +import comfy.model_management as mm +device = mm.get_torch_device() class DDPM(pl.LightningModule): # classic DDPM with Gaussian diffusion, in image space @@ -220,6 +222,9 @@ class DDPM(pl.LightningModule): variance = extract_into_tensor(1.0 - self.alphas_cumprod, t, x_start.shape) log_variance = extract_into_tensor(self.log_one_minus_alphas_cumprod, t, x_start.shape) return mean, variance, log_variance + + def get_sqrt_alpha_t_bar(self,x_start,t): + return extract_into_tensor(self.sqrt_alphas_cumprod, t, x_start.shape) def predict_start_from_noise(self, x_t, t, noise): return ( @@ -382,6 +387,7 @@ class LatentDiffusion(DDPM): logdir=None, rand_cond_frame=False, en_and_decode_n_samples_a_time=None, + control_scale=1.0, *args, **kwargs): self.num_timesteps_cond = default(num_timesteps_cond, 1) self.scale_by_std = scale_by_std @@ -399,6 +405,7 @@ class LatentDiffusion(DDPM): self.loop_video = loop_video self.fps_condition_type = fps_condition_type self.perframe_ae = perframe_ae + self.control_scale = control_scale self.logdir = logdir self.rand_cond_frame = rand_cond_frame @@ -528,17 +535,17 @@ class LatentDiffusion(DDPM): n_samples = default(self.en_and_decode_n_samples_a_time, self.temporal_length) n_rounds = math.ceil(z.shape[0] / n_samples) - with torch.autocast("cuda", enabled=True): - for n in range(n_rounds): - if isinstance(self.first_stage_model.decoder, VideoDecoder): - kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])}) - else: - kwargs = {} - - out = self.first_stage_model.decode( - z[n * n_samples : (n + 1) * n_samples], **kwargs - ) - results.append(out) + #with torch.autocast(mm.get_autocast_device(device), enabled=True): + for n in range(n_rounds): + if isinstance(self.first_stage_model.decoder, VideoDecoder): + kwargs.update({"timesteps": len(z[n * n_samples : (n + 1) * n_samples])}) + else: + kwargs = {} + + out = self.first_stage_model.decode( + z[n * n_samples : (n + 1) * n_samples], **kwargs + ) + results.append(out) results = torch.cat(results, dim=0) if reshape_back: @@ -567,9 +574,20 @@ class LatentDiffusion(DDPM): if not isinstance(cond, list): cond = [cond] key = 'c_concat' if self.model.conditioning_key == 'concat' else 'c_crossattn' - cond = {key: cond} + cond = {key: [cond[0]]} - x_recon = self.model(x_noisy, t, **cond, **kwargs) + control_cond = cond["control_cond"] + + if control_cond is not None: + control_cond = rearrange(control_cond, 'b c t h w-> (b t) c h w') + control_x = rearrange(x_noisy, 'b c t h w-> (b t) c h w') + control_context = repeat(cond["c_crossattn"][0], "b c l-> (repeat b) c l", repeat=16) + control = self.control_model(x=control_x, hint=control_cond, timesteps=t, context=control_context) + control = [c * self.control_model.control_scale for c in control] + else: + control = None + + x_recon = self.model(x_noisy, t, c_crossattn=cond["c_crossattn"], c_concat=cond["c_concat"], control=control, **kwargs) if isinstance(x_recon, tuple): return x_recon[0] @@ -699,9 +717,9 @@ class LatentDiffusion(DDPM): class LatentVisualDiffusion(LatentDiffusion): def __init__(self, img_cond_stage_config, image_proj_stage_config, freeze_embedder=True, *args, **kwargs): super().__init__(*args, **kwargs) - #self._init_embedder(img_cond_stage_config, freeze_embedder) + self._init_embedder(img_cond_stage_config, freeze_embedder) self.image_proj_model = instantiate_from_config(image_proj_stage_config) - self.embedder = None + def _init_embedder(self, config, freeze=True): embedder = instantiate_from_config(config) if freeze: @@ -717,7 +735,7 @@ class DiffusionWrapper(pl.LightningModule): self.diffusion_model = instantiate_from_config(diff_model_config) self.conditioning_key = conditioning_key - def forward(self, x, t, c_concat: list = None, c_crossattn: list = None, + def forward(self, x, t, c_concat: list = None, c_crossattn: list = None, control = None, c_adm=None, s=None, mask=None, **kwargs): # temporal_context = fps is foNone if self.conditioning_key is None: @@ -732,7 +750,7 @@ class DiffusionWrapper(pl.LightningModule): ## it is just right [b,c,t,h,w]: concatenate in channel dim xc = torch.cat([x] + c_concat, dim=1) cc = torch.cat(c_crossattn, 1) - out = self.diffusion_model(xc, t, context=cc, **kwargs) + out = self.diffusion_model(xc, t, context=cc, control=control, **kwargs) elif self.conditioning_key == 'resblockcond': cc = c_crossattn[0] out = self.diffusion_model(x, t, context=cc) diff --git a/lvdm/models/samplers/ddim.py b/lvdm/models/samplers/ddim.py index 7e8237a..ef18285 100644 --- a/lvdm/models/samplers/ddim.py +++ b/lvdm/models/samplers/ddim.py @@ -4,8 +4,10 @@ import torch from ....lvdm.models.utils_diffusion import make_ddim_sampling_parameters, make_ddim_timesteps, rescale_noise_cfg from ....lvdm.common import noise_like from ....lvdm.common import extract_into_tensor -import copy import comfy.utils +import comfy.model_management as mm + +device = mm.get_torch_device() class DDIMSampler(object): def __init__(self, model, schedule="linear", **kwargs): @@ -17,13 +19,16 @@ class DDIMSampler(object): def register_buffer(self, name, attr): if type(attr) == torch.Tensor: - if attr.device != torch.device("cuda"): - attr = attr.to(torch.device("cuda")) + if attr.device != torch.device(device): + if mm.is_device_mps(device): + attr = attr.to(torch.device(device), torch.float32) + else: + attr = attr.to(torch.device(device)) setattr(self, name, attr) - def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., verbose=True): + def make_schedule(self, ddim_num_steps, ddim_discretize="uniform", ddim_eta=0., ddpm_from=1000, verbose=True): self.ddim_timesteps = make_ddim_timesteps(ddim_discr_method=ddim_discretize, num_ddim_timesteps=ddim_num_steps, - num_ddpm_timesteps=self.ddpm_num_timesteps,verbose=verbose) + num_ddpm_timesteps=ddpm_from,verbose=verbose) alphas_cumprod = self.model.alphas_cumprod assert alphas_cumprod.shape[0] == self.ddpm_num_timesteps, 'alphas have to be defined for each timestep' to_torch = lambda x: x.clone().detach().to(torch.float32).to(self.model.device) @@ -83,6 +88,7 @@ class DDIMSampler(object): fs=None, timestep_spacing='uniform', #uniform_trailing for starting from last timestep guidance_rescale=0.0, + ddpm_from=1000, **kwargs ): @@ -100,7 +106,7 @@ class DDIMSampler(object): if conditioning.shape[0] != batch_size: print(f"Warning: Got {conditioning.shape[0]} conditionings but batch-size is {batch_size}") - self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, verbose=schedule_verbose) + self.make_schedule(ddim_num_steps=S, ddim_discretize=timestep_spacing, ddim_eta=eta, ddpm_from=ddpm_from, verbose=schedule_verbose) # make shape if len(shape) == 3: @@ -142,8 +148,10 @@ class DDIMSampler(object): device = self.model.betas.device b = shape[0] if x_T is None: + print("Using random noise") img = torch.randn(shape, device=device) else: + print("Using input noise") img = x_T if precision is not None: if precision == 16: @@ -165,6 +173,9 @@ class DDIMSampler(object): clean_cond = kwargs.pop("clean_cond", False) + sigmas = self.ddim_sigmas_for_original_num_steps if ddim_use_original_steps else self.ddim_sigmas + print("Sigmas:", sigmas) + # cond_copy, unconditional_conditioning_copy = copy.deepcopy(cond), copy.deepcopy(unconditional_conditioning) pbar = comfy.utils.ProgressBar(total_steps) for i, step in enumerate(iterator): @@ -180,10 +191,7 @@ class DDIMSampler(object): img_orig = self.model.q_sample(x0, ts) # TODO: deterministic forward pass? img = img_orig * mask + (1. - mask) * img # keep original & modify use img - - - - outs = self.p_sample_ddim(img, cond, ts, index=index, use_original_steps=ddim_use_original_steps, + outs = self.p_sample_ddim(img, cond, ts, sigmas, index=index, use_original_steps=ddim_use_original_steps, quantize_denoised=quantize_denoised, temperature=temperature, noise_dropout=noise_dropout, score_corrector=score_corrector, corrector_kwargs=corrector_kwargs, @@ -204,7 +212,7 @@ class DDIMSampler(object): return img, intermediates @torch.no_grad() - def p_sample_ddim(self, x, c, t, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False, + def p_sample_ddim(self, x, c, t, sigmas, index, repeat_noise=False, use_original_steps=False, quantize_denoised=False, temperature=1., noise_dropout=0., score_corrector=None, corrector_kwargs=None, unconditional_guidance_scale=1., unconditional_conditioning=None, uc_type=None, conditional_guidance_scale_temporal=None,mask=None,x0=None,guidance_rescale=0.0,**kwargs): @@ -242,7 +250,7 @@ class DDIMSampler(object): alphas_prev = self.model.alphas_cumprod_prev if use_original_steps else self.ddim_alphas_prev sqrt_one_minus_alphas = self.model.sqrt_one_minus_alphas_cumprod if use_original_steps else self.ddim_sqrt_one_minus_alphas # sigmas = self.model.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas - sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas + #sigmas = self.ddim_sigmas_for_original_num_steps if use_original_steps else self.ddim_sigmas # select parameters corresponding to the currently considered timestep if is_video: diff --git a/lvdm/models/utils_diffusion.py b/lvdm/models/utils_diffusion.py index 30043d3..db02a80 100644 --- a/lvdm/models/utils_diffusion.py +++ b/lvdm/models/utils_diffusion.py @@ -4,6 +4,8 @@ import torch import torch.nn.functional as F from einops import repeat +import comfy.model_management as mm +device = mm.get_torch_device() def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False): """ @@ -29,14 +31,19 @@ def timestep_embedding(timesteps, dim, max_period=10000, repeat_only=False): def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, cosine_s=8e-3): + if mm.is_device_mps(device): + dtype = torch.float32 + else: + dtype = torch.float64 + if schedule == "linear": betas = ( - torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=torch.float64) ** 2 + torch.linspace(linear_start ** 0.5, linear_end ** 0.5, n_timestep, dtype=dtype) ** 2 ) elif schedule == "cosine": timesteps = ( - torch.arange(n_timestep + 1, dtype=torch.float64) / n_timestep + cosine_s + torch.arange(n_timestep + 1, dtype=dtype) / n_timestep + cosine_s ) alphas = timesteps / (1 + cosine_s) * np.pi / 2 alphas = torch.cos(alphas).pow(2) @@ -45,9 +52,9 @@ def make_beta_schedule(schedule, n_timestep, linear_start=1e-4, linear_end=2e-2, betas = np.clip(betas, a_min=0, a_max=0.999) elif schedule == "sqrt_linear": - betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) + betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=dtype) elif schedule == "sqrt": - betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=torch.float64) ** 0.5 + betas = torch.linspace(linear_start, linear_end, n_timestep, dtype=dtype) ** 0.5 else: raise ValueError(f"schedule '{schedule}' unknown.") return betas.numpy() diff --git a/lvdm/modules/attention.py b/lvdm/modules/attention.py index 4a86de0..02889e6 100644 --- a/lvdm/modules/attention.py +++ b/lvdm/modules/attention.py @@ -16,13 +16,8 @@ from ...lvdm.common import ( ) from ...lvdm.basics import zero_module -class Conv2d(torch.nn.Conv2d): - def reset_parameters(self): - return None - -class Linear(torch.nn.Linear): - def reset_parameters(self): - return None +import comfy.ops +ops = comfy.ops.manual_cast class RelativePosition(nn.Module): """ https://github.com/evelinehong/Transformer_Relative_Position_PyTorch/blob/master/relative_position.py """ @@ -57,11 +52,11 @@ class CrossAttention(nn.Module): self.scale = dim_head**-0.5 self.heads = heads self.dim_head = dim_head - self.to_q = Linear(query_dim, inner_dim, bias=False) - self.to_k = Linear(context_dim, inner_dim, bias=False) - self.to_v = Linear(context_dim, inner_dim, bias=False) + self.to_q = ops.Linear(query_dim, inner_dim, bias=False) + self.to_k = ops.Linear(context_dim, inner_dim, bias=False) + self.to_v = ops.Linear(context_dim, inner_dim, bias=False) - self.to_out = nn.Sequential(Linear(inner_dim, query_dim), nn.Dropout(dropout)) + self.to_out = nn.Sequential(ops.Linear(inner_dim, query_dim), nn.Dropout(dropout)) self.relative_position = relative_position if self.relative_position: @@ -79,8 +74,8 @@ class CrossAttention(nn.Module): self.text_context_len = text_context_len self.image_cross_attention_scale_learnable = image_cross_attention_scale_learnable if self.image_cross_attention: - self.to_k_ip = Linear(context_dim, inner_dim, bias=False) - self.to_v_ip = Linear(context_dim, inner_dim, bias=False) + self.to_k_ip = ops.Linear(context_dim, inner_dim, bias=False) + self.to_v_ip = ops.Linear(context_dim, inner_dim, bias=False) if image_cross_attention_scale_learnable: self.register_parameter('alpha', nn.Parameter(torch.tensor(0.)) ) @@ -229,9 +224,9 @@ class BasicTransformerBlock(nn.Module): self.attn2 = attn_cls(query_dim=dim, context_dim=context_dim, heads=n_heads, dim_head=d_head, dropout=dropout, video_length=video_length, image_cross_attention=image_cross_attention, image_cross_attention_scale=image_cross_attention_scale, image_cross_attention_scale_learnable=image_cross_attention_scale_learnable,text_context_len=text_context_len) self.image_cross_attention = image_cross_attention - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) - self.norm3 = nn.LayerNorm(dim) + self.norm1 = ops.LayerNorm(dim) + self.norm2 = ops.LayerNorm(dim) + self.norm3 = ops.LayerNorm(dim) self.checkpoint = checkpoint @@ -269,11 +264,11 @@ class SpatialTransformer(nn.Module): super().__init__() self.in_channels = in_channels inner_dim = n_heads * d_head - self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + self.norm = ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) if not use_linear: - self.proj_in = Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) + self.proj_in = ops.Conv2d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) else: - self.proj_in = Linear(in_channels, inner_dim) + self.proj_in = ops.Linear(in_channels, inner_dim) attention_cls = None self.transformer_blocks = nn.ModuleList([ @@ -292,9 +287,9 @@ class SpatialTransformer(nn.Module): ) for d in range(depth) ]) if not use_linear: - self.proj_out = zero_module(Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)) + self.proj_out = zero_module(ops.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)) else: - self.proj_out = zero_module(Linear(inner_dim, in_channels)) + self.proj_out = zero_module(ops.Linear(inner_dim, in_channels)) self.use_linear = use_linear @@ -335,12 +330,12 @@ class TemporalTransformer(nn.Module): self.in_channels = in_channels inner_dim = n_heads * d_head - self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) + self.norm = ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + self.proj_in = ops.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) if not use_linear: - self.proj_in = nn.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) + self.proj_in = ops.Conv1d(in_channels, inner_dim, kernel_size=1, stride=1, padding=0) else: - self.proj_in = Linear(in_channels, inner_dim) + self.proj_in = ops.Linear(in_channels, inner_dim) if relative_position: assert(temporal_length is not None) @@ -364,9 +359,9 @@ class TemporalTransformer(nn.Module): checkpoint=use_checkpoint) for d in range(depth) ]) if not use_linear: - self.proj_out = zero_module(nn.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)) + self.proj_out = zero_module(ops.Conv1d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0)) else: - self.proj_out = zero_module(Linear(inner_dim, in_channels)) + self.proj_out = zero_module(ops.Linear(inner_dim, in_channels)) self.use_linear = use_linear @@ -444,7 +439,7 @@ class TemporalTransformer(nn.Module): class GEGLU(nn.Module): def __init__(self, dim_in, dim_out): super().__init__() - self.proj = Linear(dim_in, dim_out * 2) + self.proj = ops.Linear(dim_in, dim_out * 2) def forward(self, x): x, gate = self.proj(x).chunk(2, dim=-1) @@ -457,14 +452,14 @@ class FeedForward(nn.Module): inner_dim = int(dim * mult) dim_out = default(dim_out, dim) project_in = nn.Sequential( - Linear(dim, inner_dim), + ops.Linear(dim, inner_dim), nn.GELU() ) if not glu else GEGLU(dim, inner_dim) self.net = nn.Sequential( project_in, nn.Dropout(dropout), - Linear(inner_dim, dim_out) + ops.Linear(inner_dim, dim_out) ) def forward(self, x): @@ -476,8 +471,8 @@ class LinearAttention(nn.Module): super().__init__() self.heads = heads hidden_dim = dim_head * heads - self.to_qkv = Conv2d(dim, hidden_dim * 3, 1, bias = False) - self.to_out = Conv2d(hidden_dim, dim, 1) + self.to_qkv = ops.Conv2d(dim, hidden_dim * 3, 1, bias = False) + self.to_out = ops.Conv2d(hidden_dim, dim, 1) def forward(self, x): b, c, h, w = x.shape @@ -495,23 +490,23 @@ class SpatialSelfAttention(nn.Module): super().__init__() self.in_channels = in_channels - self.norm = torch.nn.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) - self.q = torch.Conv2d(in_channels, + self.norm = ops.GroupNorm(num_groups=32, num_channels=in_channels, eps=1e-6, affine=True) + self.q = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = torch.Conv2d(in_channels, + self.k = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = torch.Conv2d(in_channels, + self.v = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = torch.Conv2d(in_channels, + self.proj_out = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, diff --git a/lvdm/modules/attention_svd.py b/lvdm/modules/attention_svd.py index 6c5852c..d2fbaeb 100644 --- a/lvdm/modules/attention_svd.py +++ b/lvdm/modules/attention_svd.py @@ -12,6 +12,9 @@ from torch.utils.checkpoint import checkpoint logpy = logging.getLogger(__name__) +import comfy.ops +ops = comfy.ops.manual_cast + if version.parse(torch.__version__) >= version.parse("2.0.0"): SDP_IS_AVAILABLE = True from torch.backends.cuda import SDPBackend, sdp_kernel @@ -87,7 +90,7 @@ def init_(tensor): class GEGLU(nn.Module): def __init__(self, dim_in, dim_out): super().__init__() - self.proj = nn.Linear(dim_in, dim_out * 2) + self.proj = ops.Linear(dim_in, dim_out * 2) def forward(self, x): x, gate = self.proj(x).chunk(2, dim=-1) @@ -100,13 +103,13 @@ class FeedForward(nn.Module): inner_dim = int(dim * mult) dim_out = default(dim_out, dim) project_in = ( - nn.Sequential(nn.Linear(dim, inner_dim), nn.GELU()) + nn.Sequential(ops.Linear(dim, inner_dim), nn.GELU()) if not glu else GEGLU(dim, inner_dim) ) self.net = nn.Sequential( - project_in, nn.Dropout(dropout), nn.Linear(inner_dim, dim_out) + project_in, nn.Dropout(dropout), ops.Linear(inner_dim, dim_out) ) def forward(self, x): @@ -123,7 +126,7 @@ def zero_module(module): def Normalize(in_channels): - return torch.nn.GroupNorm( + return ops.GroupNorm( num_groups=32, num_channels=in_channels, eps=1e-6, affine=True ) @@ -133,8 +136,8 @@ class LinearAttention(nn.Module): super().__init__() self.heads = heads hidden_dim = dim_head * heads - self.to_qkv = nn.Conv2d(dim, hidden_dim * 3, 1, bias=False) - self.to_out = nn.Conv2d(hidden_dim, dim, 1) + self.to_qkv = ops.Conv2d(dim, hidden_dim * 3, 1, bias=False) + self.to_out = ops.Conv2d(hidden_dim, dim, 1) def forward(self, x): b, c, h, w = x.shape @@ -169,9 +172,9 @@ class SelfAttention(nn.Module): head_dim = dim // num_heads self.scale = qk_scale or head_dim**-0.5 - self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) + self.qkv = ops.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) - self.proj = nn.Linear(dim, dim) + self.proj = ops.Linear(dim, dim) self.proj_drop = nn.Dropout(proj_drop) assert attn_mode in self.ATTENTION_MODES self.attn_mode = attn_mode @@ -213,16 +216,16 @@ class SpatialSelfAttention(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d( + self.q = torch.ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.k = torch.nn.Conv2d( + self.k = torch.ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.v = torch.nn.Conv2d( + self.v = torch.ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) - self.proj_out = torch.nn.Conv2d( + self.proj_out = torch.ops.Conv2d( in_channels, in_channels, kernel_size=1, stride=1, padding=0 ) @@ -269,12 +272,12 @@ class CrossAttention(nn.Module): self.scale = dim_head**-0.5 self.heads = heads - self.to_q = nn.Linear(query_dim, inner_dim, bias=False) - self.to_k = nn.Linear(context_dim, inner_dim, bias=False) - self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_q = ops.Linear(query_dim, inner_dim, bias=False) + self.to_k = ops.Linear(context_dim, inner_dim, bias=False) + self.to_v = ops.Linear(context_dim, inner_dim, bias=False) self.to_out = nn.Sequential( - nn.Linear(inner_dim, query_dim), nn.Dropout(dropout) + ops.Linear(inner_dim, query_dim), nn.Dropout(dropout) ) self.backend = backend @@ -361,12 +364,12 @@ class MemoryEfficientCrossAttention(nn.Module): self.heads = heads self.dim_head = dim_head - self.to_q = nn.Linear(query_dim, inner_dim, bias=False) - self.to_k = nn.Linear(context_dim, inner_dim, bias=False) - self.to_v = nn.Linear(context_dim, inner_dim, bias=False) + self.to_q = ops.Linear(query_dim, inner_dim, bias=False) + self.to_k = ops.Linear(context_dim, inner_dim, bias=False) + self.to_v = ops.Linear(context_dim, inner_dim, bias=False) self.to_out = nn.Sequential( - nn.Linear(inner_dim, query_dim), nn.Dropout(dropout) + ops.Linear(inner_dim, query_dim), nn.Dropout(dropout) ) self.attention_op: Optional[Any] = None @@ -517,9 +520,9 @@ class BasicTransformerBlock(nn.Module): dropout=dropout, backend=sdp_backend, ) # is self-attn if context is none - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) - self.norm3 = nn.LayerNorm(dim) + self.norm1 = ops.LayerNorm(dim) + self.norm2 = ops.LayerNorm(dim) + self.norm3 = ops.LayerNorm(dim) self.checkpoint = checkpoint if self.checkpoint: logpy.debug(f"{self.__class__.__name__} is using checkpointing") @@ -601,8 +604,8 @@ class BasicTransformerSingleLayerBlock(nn.Module): context_dim=context_dim, ) self.ff = FeedForward(dim, dropout=dropout, glu=gated_ff) - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) + self.norm1 = ops.LayerNorm(dim) + self.norm2 = ops.LayerNorm(dim) self.checkpoint = checkpoint def forward(self, x, context=None): @@ -668,11 +671,11 @@ class SpatialTransformer(nn.Module): inner_dim = n_heads * d_head self.norm = Normalize(in_channels) if not use_linear: - self.proj_in = nn.Conv2d( + self.proj_in = ops.Conv2d( in_channels, inner_dim, kernel_size=1, stride=1, padding=0 ) else: - self.proj_in = nn.Linear(in_channels, inner_dim) + self.proj_in = ops.Linear(in_channels, inner_dim) self.transformer_blocks = nn.ModuleList( [ @@ -692,11 +695,11 @@ class SpatialTransformer(nn.Module): ) if not use_linear: self.proj_out = zero_module( - nn.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0) + ops.Conv2d(inner_dim, in_channels, kernel_size=1, stride=1, padding=0) ) else: # self.proj_out = zero_module(nn.Linear(in_channels, inner_dim)) - self.proj_out = zero_module(nn.Linear(inner_dim, in_channels)) + self.proj_out = zero_module(ops.Linear(inner_dim, in_channels)) self.use_linear = use_linear def forward(self, x, context=None): diff --git a/lvdm/modules/encoders/condition.py b/lvdm/modules/encoders/condition.py index d2a55ce..b2dd4c0 100644 --- a/lvdm/modules/encoders/condition.py +++ b/lvdm/modules/encoders/condition.py @@ -1,12 +1,14 @@ import torch import torch.nn as nn -import kornia -import open_clip -from torch.utils.checkpoint import checkpoint -from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel +#import kornia +#import open_clip +#from torch.utils.checkpoint import checkpoint +#from transformers import T5Tokenizer, T5EncoderModel, CLIPTokenizer, CLIPTextModel from ....lvdm.common import autocast from ....utils.utils import count_params +import comfy.model_management as mm +device = mm.get_torch_device() class AbstractEncoder(nn.Module): def __init__(self): @@ -41,7 +43,7 @@ class ClassEmbedder(nn.Module): c = self.embedding(c) return c - def get_unconditional_conditioning(self, bs, device="cuda"): + def get_unconditional_conditioning(self, bs, device=device): uc_class = self.n_classes - 1 # 1000 classes --> 0 ... 999, one extra class for ucg (class 1000) uc = torch.ones((bs,), device=device) * uc_class uc = {self.key: uc} @@ -57,7 +59,7 @@ def disabled_train(self, mode=True): class FrozenT5Embedder(AbstractEncoder): """Uses the T5 transformer encoder for text""" - def __init__(self, version="google/t5-v1_1-large", device="cuda", max_length=77, + def __init__(self, version="google/t5-v1_1-large", device=device, max_length=77, freeze=True): # others are google/t5-v1_1-xl and google/t5-v1_1-xxl super().__init__() self.tokenizer = T5Tokenizer.from_pretrained(version) @@ -94,7 +96,7 @@ class FrozenCLIPEmbedder(AbstractEncoder): "hidden" ] - def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77, + def __init__(self, version="openai/clip-vit-large-patch14", device=device, max_length=77, freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32 super().__init__() assert layer in self.LAYERS @@ -138,7 +140,7 @@ class ClipImageEmbedder(nn.Module): self, model, jit=False, - device='cuda' if torch.cuda.is_available() else 'cpu', + device=device, antialias=True, ucg_rate=0. ): @@ -181,7 +183,7 @@ class FrozenOpenCLIPEmbedder(AbstractEncoder): "penultimate" ] - def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77, + def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, max_length=77, freeze=True, layer="last"): super().__init__() assert layer in self.LAYERS @@ -239,7 +241,7 @@ class FrozenOpenCLIPImageEmbedder(AbstractEncoder): Uses the OpenCLIP vision transformer encoder for images """ - def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", max_length=77, + def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, max_length=77, freeze=True, layer="pooled", antialias=True, ucg_rate=0.): super().__init__() model, _, _ = open_clip.create_model_and_transforms(arch, device=torch.device('cpu'), @@ -297,7 +299,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder): Uses the OpenCLIP vision transformer encoder for images """ - def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device="cuda", + def __init__(self, arch="ViT-H-14", version="laion2b_s32b_b79k", device=device, freeze=True, layer="pooled", antialias=True): super().__init__() return @@ -373,7 +375,7 @@ class FrozenOpenCLIPImageEmbedderV2(AbstractEncoder): return x class FrozenCLIPT5Encoder(AbstractEncoder): - def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device="cuda", + def __init__(self, clip_version="openai/clip-vit-large-patch14", t5_version="google/t5-v1_1-xl", device=device, clip_max_length=77, t5_max_length=77): super().__init__() self.clip_encoder = FrozenCLIPEmbedder(clip_version, device, max_length=clip_max_length) diff --git a/lvdm/modules/encoders/resampler.py b/lvdm/modules/encoders/resampler.py index 0c30c58..6a0d27e 100644 --- a/lvdm/modules/encoders/resampler.py +++ b/lvdm/modules/encoders/resampler.py @@ -5,6 +5,8 @@ import math import torch import torch.nn as nn +import comfy.ops +ops = comfy.ops.manual_cast class ImageProjModel(nn.Module): """Projection Model""" @@ -12,8 +14,8 @@ class ImageProjModel(nn.Module): super().__init__() self.cross_attention_dim = cross_attention_dim self.clip_extra_context_tokens = clip_extra_context_tokens - self.proj = nn.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim) - self.norm = nn.LayerNorm(cross_attention_dim) + self.proj = ops.Linear(clip_embeddings_dim, self.clip_extra_context_tokens * cross_attention_dim) + self.norm = ops.LayerNorm(cross_attention_dim) def forward(self, image_embeds): #embeds = image_embeds @@ -27,10 +29,10 @@ class ImageProjModel(nn.Module): def FeedForward(dim, mult=4): inner_dim = int(dim * mult) return nn.Sequential( - nn.LayerNorm(dim), - nn.Linear(dim, inner_dim, bias=False), + ops.LayerNorm(dim), + ops.Linear(dim, inner_dim, bias=False), nn.GELU(), - nn.Linear(inner_dim, dim, bias=False), + ops.Linear(inner_dim, dim, bias=False), ) @@ -53,12 +55,12 @@ class PerceiverAttention(nn.Module): self.heads = heads inner_dim = dim_head * heads - self.norm1 = nn.LayerNorm(dim) - self.norm2 = nn.LayerNorm(dim) + self.norm1 = ops.LayerNorm(dim) + self.norm2 = ops.LayerNorm(dim) - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_kv = nn.Linear(dim, inner_dim * 2, bias=False) - self.to_out = nn.Linear(inner_dim, dim, bias=False) + self.to_q = ops.Linear(dim, inner_dim, bias=False) + self.to_kv = ops.Linear(dim, inner_dim * 2, bias=False) + self.to_out = ops.Linear(inner_dim, dim, bias=False) def forward(self, x, latents): @@ -116,9 +118,9 @@ class Resampler(nn.Module): num_queries = num_queries * video_length self.latents = nn.Parameter(torch.randn(1, num_queries, dim) / dim**0.5) - self.proj_in = nn.Linear(embedding_dim, dim) - self.proj_out = nn.Linear(dim, output_dim) - self.norm_out = nn.LayerNorm(output_dim) + self.proj_in = ops.Linear(embedding_dim, dim) + self.proj_out = ops.Linear(dim, output_dim) + self.norm_out = ops.LayerNorm(output_dim) self.layers = nn.ModuleList([]) for _ in range(depth): diff --git a/lvdm/modules/networks/ae_modules.py b/lvdm/modules/networks/ae_modules.py index 55f0e4f..08b8e6a 100644 --- a/lvdm/modules/networks/ae_modules.py +++ b/lvdm/modules/networks/ae_modules.py @@ -7,13 +7,16 @@ from einops import rearrange from ....utils.utils import instantiate_from_config from ....lvdm.modules.attention import LinearAttention +import comfy.ops +ops = comfy.ops.manual_cast + def nonlinearity(x): # swish return x*torch.sigmoid(x) def Normalize(in_channels, num_groups=32): - return torch.nn.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True) + return ops.GroupNorm(num_groups=num_groups, num_channels=in_channels, eps=1e-6, affine=True) @@ -29,22 +32,22 @@ class AttnBlock(nn.Module): self.in_channels = in_channels self.norm = Normalize(in_channels) - self.q = torch.nn.Conv2d(in_channels, + self.q = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.k = torch.nn.Conv2d(in_channels, + self.k = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.v = torch.nn.Conv2d(in_channels, + self.v = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, padding=0) - self.proj_out = torch.nn.Conv2d(in_channels, + self.proj_out = ops.Conv2d(in_channels, in_channels, kernel_size=1, stride=1, @@ -94,7 +97,7 @@ class Downsample(nn.Module): self.in_channels = in_channels if self.with_conv: # no asymmetric padding in torch conv, must do it ourselves - self.conv = torch.nn.Conv2d(in_channels, + self.conv = ops.Conv2d(in_channels, in_channels, kernel_size=3, stride=2, @@ -114,7 +117,7 @@ class Upsample(nn.Module): self.with_conv = with_conv self.in_channels = in_channels if self.with_conv: - self.conv = torch.nn.Conv2d(in_channels, + self.conv = ops.Conv2d(in_channels, in_channels, kernel_size=3, stride=1, @@ -158,30 +161,30 @@ class ResnetBlock(nn.Module): self.use_conv_shortcut = conv_shortcut self.norm1 = Normalize(in_channels) - self.conv1 = torch.nn.Conv2d(in_channels, + self.conv1 = ops.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) if temb_channels > 0: - self.temb_proj = torch.nn.Linear(temb_channels, + self.temb_proj = ops.Linear(temb_channels, out_channels) self.norm2 = Normalize(out_channels) self.dropout = torch.nn.Dropout(dropout) - self.conv2 = torch.nn.Conv2d(out_channels, + self.conv2 = ops.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1) if self.in_channels != self.out_channels: if self.use_conv_shortcut: - self.conv_shortcut = torch.nn.Conv2d(in_channels, + self.conv_shortcut = ops.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, padding=1) else: - self.nin_shortcut = torch.nn.Conv2d(in_channels, + self.nin_shortcut = ops.Conv2d(in_channels, out_channels, kernel_size=1, stride=1, @@ -227,14 +230,14 @@ class Model(nn.Module): # timestep embedding self.temb = nn.Module() self.temb.dense = nn.ModuleList([ - torch.nn.Linear(self.ch, + ops.Linear(self.ch, self.temb_ch), - torch.nn.Linear(self.temb_ch, + ops.Linear(self.temb_ch, self.temb_ch), ]) # downsampling - self.conv_in = torch.nn.Conv2d(in_channels, + self.conv_in = ops.Conv2d(in_channels, self.ch, kernel_size=3, stride=1, @@ -303,7 +306,7 @@ class Model(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, out_ch, kernel_size=3, stride=1, @@ -376,7 +379,7 @@ class Encoder(nn.Module): self.in_channels = in_channels # downsampling - self.conv_in = torch.nn.Conv2d(in_channels, + self.conv_in = ops.Conv2d(in_channels, self.ch, kernel_size=3, stride=1, @@ -421,7 +424,7 @@ class Encoder(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, 2*z_channels if double_z else z_channels, kernel_size=3, stride=1, @@ -498,7 +501,7 @@ class Decoder(nn.Module): self.z_shape, np.prod(self.z_shape))) # z to block_in - self.conv_in = torch.nn.Conv2d(z_channels, + self.conv_in = ops.Conv2d(z_channels, block_in, kernel_size=3, stride=1, @@ -540,7 +543,7 @@ class Decoder(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, out_ch, kernel_size=3, stride=1, @@ -591,7 +594,7 @@ class Decoder(nn.Module): class SimpleDecoder(nn.Module): def __init__(self, in_channels, out_channels, *args, **kwargs): super().__init__() - self.model = nn.ModuleList([nn.Conv2d(in_channels, in_channels, 1), + self.model = nn.ModuleList([ops.Conv2d(in_channels, in_channels, 1), ResnetBlock(in_channels=in_channels, out_channels=2 * in_channels, temb_channels=0, dropout=0.0), @@ -601,11 +604,11 @@ class SimpleDecoder(nn.Module): ResnetBlock(in_channels=4 * in_channels, out_channels=2 * in_channels, temb_channels=0, dropout=0.0), - nn.Conv2d(2*in_channels, in_channels, 1), + ops.Conv2d(2*in_channels, in_channels, 1), Upsample(in_channels, with_conv=True)]) # end self.norm_out = Normalize(in_channels) - self.conv_out = torch.nn.Conv2d(in_channels, + self.conv_out = ops.Conv2d(in_channels, out_channels, kernel_size=3, stride=1, @@ -652,7 +655,7 @@ class UpsampleDecoder(nn.Module): # end self.norm_out = Normalize(block_in) - self.conv_out = torch.nn.Conv2d(block_in, + self.conv_out = ops.Conv2d(block_in, out_channels, kernel_size=3, stride=1, @@ -677,7 +680,7 @@ class LatentRescaler(nn.Module): super().__init__() # residual block, interpolate, residual block self.factor = factor - self.conv_in = nn.Conv2d(in_channels, + self.conv_in = ops.Conv2d(in_channels, mid_channels, kernel_size=3, stride=1, @@ -692,7 +695,7 @@ class LatentRescaler(nn.Module): temb_channels=0, dropout=0.0) for _ in range(depth)]) - self.conv_out = nn.Conv2d(mid_channels, + self.conv_out = ops.Conv2d(mid_channels, out_channels, kernel_size=1, ) @@ -774,7 +777,7 @@ class Resize(nn.Module): raise NotImplementedError() assert in_channels is not None # no asymmetric padding in torch conv, must do it ourselves - self.conv = torch.nn.Conv2d(in_channels, + self.conv = ops.Conv2d(in_channels, in_channels, kernel_size=4, stride=2, @@ -809,7 +812,7 @@ class FirstStagePostProcessor(nn.Module): n_channels = self.pretrained_model.encoder.ch self.proj_norm = Normalize(in_channels,num_groups=in_channels//2) - self.proj = nn.Conv2d(in_channels,n_channels,kernel_size=3, + self.proj = ops.Conv2d(in_channels,n_channels,kernel_size=3, stride=1,padding=1) blocks = [] diff --git a/lvdm/modules/networks/openaimodel3d.py b/lvdm/modules/networks/openaimodel3d.py index a1b29b8..1308174 100644 --- a/lvdm/modules/networks/openaimodel3d.py +++ b/lvdm/modules/networks/openaimodel3d.py @@ -2,7 +2,9 @@ from functools import partial from abc import abstractmethod import torch import torch.nn as nn +import numpy as np from einops import rearrange +import math import torch.nn.functional as F from ....lvdm.models.utils_diffusion import timestep_embedding from ....lvdm.common import checkpoint @@ -15,6 +17,11 @@ from ....lvdm.basics import ( ) from ....lvdm.modules.attention import SpatialTransformer, TemporalTransformer +import comfy.ops +ops = comfy.ops.manual_cast + +def exists(x): + return x is not None class TimestepBlock(nn.Module): """ @@ -167,7 +174,7 @@ class ResBlock(TimestepBlock): self.emb_layers = nn.Sequential( nn.SiLU(), - nn.Linear( + ops.Linear( emb_channels, 2 * self.out_channels if use_scale_shift_norm else self.out_channels, ), @@ -176,7 +183,7 @@ class ResBlock(TimestepBlock): normalization(self.out_channels), nn.SiLU(), nn.Dropout(p=dropout), - zero_module(nn.Conv2d(self.out_channels, self.out_channels, 3, padding=1)), + zero_module(ops.Conv2d(self.out_channels, self.out_channels, 3, padding=1)), ) if self.out_channels == channels: @@ -253,17 +260,17 @@ class TemporalConvBlock(nn.Module): # conv layers self.conv1 = nn.Sequential( - nn.GroupNorm(32, in_channels), nn.SiLU(), - nn.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape)) + ops.GroupNorm(32, in_channels), nn.SiLU(), + ops.Conv3d(in_channels, out_channels, th_kernel_shape, padding=th_padding_shape)) self.conv2 = nn.Sequential( - nn.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout), - nn.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape)) + ops.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout), + ops.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape)) self.conv3 = nn.Sequential( - nn.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout), - nn.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape)) + ops.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout), + ops.Conv3d(out_channels, in_channels, th_kernel_shape, padding=th_padding_shape)) self.conv4 = nn.Sequential( - nn.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout), - nn.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape)) + ops.GroupNorm(32, out_channels), nn.SiLU(), nn.Dropout(dropout), + ops.Conv3d(out_channels, in_channels, tw_kernel_shape, padding=tw_padding_shape)) # zero out the last layer params,so the conv block is identity nn.init.zeros_(self.conv4[-1].weight) @@ -545,7 +552,7 @@ class UNetModel(nn.Module): zero_module(conv_nd(dims, model_channels, out_channels, 3, padding=1)), ) - def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, **kwargs): + def forward(self, x, timesteps, context=None, features_adapter=None, fs=None, frame_window_size=None, frame_window_stride=None, control=None, **kwargs): b,_,t,_,_ = x.shape t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False).type(x.dtype) emb = self.time_embed(t_emb) @@ -592,8 +599,15 @@ class UNetModel(nn.Module): assert len(features_adapter)==adapter_idx, 'Wrong features_adapter' h = self.middle_block(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) + + if control is not None: + h += control.pop() + for module in self.output_blocks: - h = torch.cat([h, hs.pop()], dim=1) + if control is None: + h = torch.cat([h, hs.pop()], dim=1) + else: + h = torch.cat([h, hs.pop() + control.pop()], dim=1) h = module(h, emb, context=context, batch_size=b, frame_window_size=frame_window_size, frame_window_stride=frame_window_stride) h = h.type(x.dtype) y = self.out(h) @@ -601,3 +615,394 @@ class UNetModel(nn.Module): # reshape back to (b c t h w) y = rearrange(y, '(b t) c h w -> b c t h w', b=b) return y + +class ControlNet(nn.Module): + def __init__( + self, + image_size, + in_channels, + model_channels, + hint_channels, + num_res_blocks, + attention_resolutions, + dropout=0, + channel_mult=(1, 2, 4, 8), + conv_resample=True, + dims=2, + use_checkpoint=False, + use_fp16=False, + num_heads=-1, + num_head_channels=-1, + num_heads_upsample=-1, + use_scale_shift_norm=False, + resblock_updown=False, + use_new_attention_order=False, + use_spatial_transformer=False, # custom transformer support + transformer_depth=1, # custom transformer support + context_dim=None, # custom transformer support + n_embed=None, # custom support for prediction of discrete ids into codebook of first stage vq model + legacy=True, + disable_self_attentions=None, + num_attention_blocks=None, + disable_middle_self_attn=False, + use_linear_in_transformer=False, + ): + super().__init__() + if use_spatial_transformer: + assert context_dim is not None, 'Fool!! You forgot to include the dimension of your cross-attention conditioning...' + + if context_dim is not None: + assert use_spatial_transformer, 'Fool!! You forgot to use the spatial transformer for your cross-attention conditioning...' + from omegaconf.listconfig import ListConfig + if type(context_dim) == ListConfig: + context_dim = list(context_dim) + + if num_heads_upsample == -1: + num_heads_upsample = num_heads + + if num_heads == -1: + assert num_head_channels != -1, 'Either num_heads or num_head_channels has to be set' + + if num_head_channels == -1: + assert num_heads != -1, 'Either num_heads or num_head_channels has to be set' + + self.dims = dims + self.image_size = image_size + self.in_channels = in_channels + self.model_channels = model_channels + if isinstance(num_res_blocks, int): + self.num_res_blocks = len(channel_mult) * [num_res_blocks] + else: + if len(num_res_blocks) != len(channel_mult): + raise ValueError("provide num_res_blocks either as an int (globally constant) or " + "as a list/tuple (per-level) with the same length as channel_mult") + self.num_res_blocks = num_res_blocks + if disable_self_attentions is not None: + # should be a list of booleans, indicating whether to disable self-attention in TransformerBlocks or not + assert len(disable_self_attentions) == len(channel_mult) + if num_attention_blocks is not None: + assert len(num_attention_blocks) == len(self.num_res_blocks) + assert all(map(lambda i: self.num_res_blocks[i] >= num_attention_blocks[i], range(len(num_attention_blocks)))) + print(f"Constructor of UNetModel received num_attention_blocks={num_attention_blocks}. " + f"This option has LESS priority than attention_resolutions {attention_resolutions}, " + f"i.e., in cases where num_attention_blocks[i] > 0 but 2**i not in attention_resolutions, " + f"attention will still not be set.") + + self.attention_resolutions = attention_resolutions + self.dropout = dropout + self.channel_mult = channel_mult + self.conv_resample = conv_resample + self.use_checkpoint = use_checkpoint + self.dtype = torch.float16 if use_fp16 else torch.float32 + self.num_heads = num_heads + self.num_head_channels = num_head_channels + self.num_heads_upsample = num_heads_upsample + self.predict_codebook_ids = n_embed is not None + + time_embed_dim = model_channels * 4 + self.time_embed = nn.Sequential( + linear(model_channels, time_embed_dim), + nn.SiLU(), + linear(time_embed_dim, time_embed_dim), + ) + + self.input_blocks = nn.ModuleList( + [ + TimestepEmbedSequential( + conv_nd(dims, in_channels, model_channels, 3, padding=1) + ) + ] + ) + self.zero_convs = nn.ModuleList([self.make_zero_conv(model_channels)]) + + self.input_hint_block = TimestepEmbedSequential( + conv_nd(dims, hint_channels, 16, 3, padding=1), + nn.SiLU(), + conv_nd(dims, 16, 16, 3, padding=1), + nn.SiLU(), + conv_nd(dims, 16, 32, 3, padding=1, stride=2), + nn.SiLU(), + conv_nd(dims, 32, 32, 3, padding=1), + nn.SiLU(), + conv_nd(dims, 32, 96, 3, padding=1, stride=2), + nn.SiLU(), + conv_nd(dims, 96, 96, 3, padding=1), + nn.SiLU(), + conv_nd(dims, 96, 256, 3, padding=1, stride=2), + nn.SiLU(), + zero_module(conv_nd(dims, 256, model_channels, 3, padding=1)) + ) + + self._feature_size = model_channels + input_block_chans = [model_channels] + ch = model_channels + ds = 1 + for level, mult in enumerate(channel_mult): + for nr in range(self.num_res_blocks[level]): + layers = [ + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=mult * model_channels, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ) + ] + ch = mult * model_channels + if ds in attention_resolutions: + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + if legacy: + # num_heads = 1 + dim_head = ch // num_heads if use_spatial_transformer else num_head_channels + if exists(disable_self_attentions): + disabled_sa = disable_self_attentions[level] + else: + disabled_sa = False + + if not exists(num_attention_blocks) or nr < num_attention_blocks[level]: + layers.append( + AttentionBlock( + ch, + use_checkpoint=use_checkpoint, + num_heads=num_heads, + num_head_channels=dim_head, + use_new_attention_order=use_new_attention_order, + ) if not use_spatial_transformer else SpatialTransformer( + ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim, + disable_self_attn=disabled_sa, use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint + ) + ) + self.input_blocks.append(TimestepEmbedSequential(*layers)) + self.zero_convs.append(self.make_zero_conv(ch)) + self._feature_size += ch + input_block_chans.append(ch) + if level != len(channel_mult) - 1: + out_ch = ch + self.input_blocks.append( + TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + out_channels=out_ch, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + down=True, + ) + if resblock_updown + else Downsample( + ch, conv_resample, dims=dims, out_channels=out_ch + ) + ) + ) + ch = out_ch + input_block_chans.append(ch) + self.zero_convs.append(self.make_zero_conv(ch)) + ds *= 2 + self._feature_size += ch + + if num_head_channels == -1: + dim_head = ch // num_heads + else: + num_heads = ch // num_head_channels + dim_head = num_head_channels + if legacy: + # num_heads = 1 + dim_head = ch // num_heads if use_spatial_transformer else num_head_channels + self.middle_block = TimestepEmbedSequential( + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + AttentionBlock( + ch, + use_checkpoint=use_checkpoint, + num_heads=num_heads, + num_head_channels=dim_head, + use_new_attention_order=use_new_attention_order, + ) if not use_spatial_transformer else SpatialTransformer( # always uses a self-attn + ch, num_heads, dim_head, depth=transformer_depth, context_dim=context_dim, + disable_self_attn=disable_middle_self_attn, use_linear=use_linear_in_transformer, + use_checkpoint=use_checkpoint + ), + ResBlock( + ch, + time_embed_dim, + dropout, + dims=dims, + use_checkpoint=use_checkpoint, + use_scale_shift_norm=use_scale_shift_norm, + ), + ) + self.middle_block_out = self.make_zero_conv(ch) + self._feature_size += ch + + def make_zero_conv(self, channels): + return TimestepEmbedSequential(zero_module(conv_nd(self.dims, channels, channels, 1, padding=0))) + + def forward(self, x, hint, timesteps, context, **kwargs): + t_emb = timestep_embedding(timesteps, self.model_channels, repeat_only=False) + emb = self.time_embed(t_emb) + + guided_hint = self.input_hint_block(hint, emb, context) + + outs = [] + + h = x.type(self.dtype) + + for module, zero_conv in zip(self.input_blocks, self.zero_convs): + if guided_hint is not None: + h = module(h, emb, context) + h += guided_hint + guided_hint = None + else: + h = module(h, emb, context) + outs.append(zero_conv(h, emb, context, True)) + + h = self.middle_block(h, emb, context) + outs.append(self.middle_block_out(h, emb, context)) + + return outs + +class AttentionBlock(nn.Module): + """ + An attention block that allows spatial positions to attend to each other. + Originally ported from here, but adapted to the N-d case. + https://github.com/hojonathanho/diffusion/blob/1e0dceb3b3495bbe19116a5e1b3596cd0706c543/diffusion_tf/models/unet.py#L66. + """ + + def __init__( + self, + channels, + num_heads=1, + num_head_channels=-1, + use_checkpoint=False, + use_new_attention_order=False, + ): + super().__init__() + self.channels = channels + if num_head_channels == -1: + self.num_heads = num_heads + else: + assert ( + channels % num_head_channels == 0 + ), f"q,k,v channels {channels} is not divisible by num_head_channels {num_head_channels}" + self.num_heads = channels // num_head_channels + self.use_checkpoint = use_checkpoint + self.norm = normalization(channels) + self.qkv = conv_nd(1, channels, channels * 3, 1) + if use_new_attention_order: + # split qkv before split heads + self.attention = QKVAttention(self.num_heads) + else: + # split heads before split qkv + self.attention = QKVAttentionLegacy(self.num_heads) + + self.proj_out = zero_module(conv_nd(1, channels, channels, 1)) + + def forward(self, x): + return checkpoint(self._forward, (x,), self.parameters(), True) # TODO: check checkpoint usage, is True # TODO: fix the .half call!!! + #return pt_checkpoint(self._forward, x) # pytorch + + def _forward(self, x): + b, c, *spatial = x.shape + x = x.reshape(b, c, -1) + qkv = self.qkv(self.norm(x)) + h = self.attention(qkv) + h = self.proj_out(h) + return (x + h).reshape(b, c, *spatial) + +class QKVAttention(nn.Module): + """ + A module which performs QKV attention and splits in a different order. + """ + + def __init__(self, n_heads): + super().__init__() + self.n_heads = n_heads + + def forward(self, qkv): + """ + Apply QKV attention. + :param qkv: an [N x (3 * H * C) x T] tensor of Qs, Ks, and Vs. + :return: an [N x (H * C) x T] tensor after attention. + """ + bs, width, length = qkv.shape + assert width % (3 * self.n_heads) == 0 + ch = width // (3 * self.n_heads) + q, k, v = qkv.chunk(3, dim=1) + scale = 1 / math.sqrt(math.sqrt(ch)) + weight = torch.einsum( + "bct,bcs->bts", + (q * scale).view(bs * self.n_heads, ch, length), + (k * scale).view(bs * self.n_heads, ch, length), + ) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + a = torch.einsum("bts,bcs->bct", weight, v.reshape(bs * self.n_heads, ch, length)) + return a.reshape(bs, -1, length) + + @staticmethod + def count_flops(model, _x, y): + return count_flops_attn(model, _x, y) + +class QKVAttentionLegacy(nn.Module): + """ + A module which performs QKV attention. Matches legacy QKVAttention + input/ouput heads shaping + """ + + def __init__(self, n_heads): + super().__init__() + self.n_heads = n_heads + + def forward(self, qkv): + """ + Apply QKV attention. + :param qkv: an [N x (H * 3 * C) x T] tensor of Qs, Ks, and Vs. + :return: an [N x (H * C) x T] tensor after attention. + """ + bs, width, length = qkv.shape + assert width % (3 * self.n_heads) == 0 + ch = width // (3 * self.n_heads) + q, k, v = qkv.reshape(bs * self.n_heads, ch * 3, length).split(ch, dim=1) + scale = 1 / math.sqrt(math.sqrt(ch)) + weight = torch.einsum( + "bct,bcs->bts", q * scale, k * scale + ) # More stable with f16 than dividing afterwards + weight = torch.softmax(weight.float(), dim=-1).type(weight.dtype) + a = torch.einsum("bts,bcs->bct", weight, v) + return a.reshape(bs, -1, length) + + @staticmethod + def count_flops(model, _x, y): + return count_flops_attn(model, _x, y) + +def count_flops_attn(model, _x, y): + """ + A counter for the `thop` package to count the operations in an + attention operation. + Meant to be used like: + macs, params = thop.profile( + model, + inputs=(inputs, timestamps), + custom_ops={QKVAttention: QKVAttention.count_flops}, + ) + """ + b, c, *spatial = y[0].shape + num_spatial = int(np.prod(spatial)) + # We perform two matmuls with the same number of ops. + # The first computes the weight matrix, the second computes + # the combination of the value vectors. + matmul_ops = 2 * b * (num_spatial ** 2) * c + model.total_ops += torch.DoubleTensor([matmul_ops]) \ No newline at end of file diff --git a/lvdm/modules/x_transformer.py b/lvdm/modules/x_transformer.py index 5321012..9757b7d 100644 --- a/lvdm/modules/x_transformer.py +++ b/lvdm/modules/x_transformer.py @@ -7,6 +7,9 @@ import torch from torch import nn, einsum import torch.nn.functional as F +import comfy.ops +ops = comfy.ops.manual_cast + # constants DEFAULT_DIM_HEAD = 64 @@ -183,7 +186,7 @@ class GRUGating(nn.Module): class GEGLU(nn.Module): def __init__(self, dim_in, dim_out): super().__init__() - self.proj = nn.Linear(dim_in, dim_out * 2) + self.proj = ops.Linear(dim_in, dim_out * 2) def forward(self, x): x, gate = self.proj(x).chunk(2, dim=-1) @@ -196,14 +199,14 @@ class FeedForward(nn.Module): inner_dim = int(dim * mult) dim_out = default(dim_out, dim) project_in = nn.Sequential( - nn.Linear(dim, inner_dim), + ops.Linear(dim, inner_dim), nn.GELU() ) if not glu else GEGLU(dim, inner_dim) self.net = nn.Sequential( project_in, nn.Dropout(dropout), - nn.Linear(inner_dim, dim_out) + ops.Linear(inner_dim, dim_out) ) def forward(self, x): @@ -236,9 +239,9 @@ class Attention(nn.Module): inner_dim = dim_head * heads - self.to_q = nn.Linear(dim, inner_dim, bias=False) - self.to_k = nn.Linear(dim, inner_dim, bias=False) - self.to_v = nn.Linear(dim, inner_dim, bias=False) + self.to_q = ops.Linear(dim, inner_dim, bias=False) + self.to_k = ops.Linear(dim, inner_dim, bias=False) + self.to_v = ops.Linear(dim, inner_dim, bias=False) self.dropout = nn.Dropout(dropout) # talking heads @@ -262,7 +265,7 @@ class Attention(nn.Module): # attention on attention self.attn_on_attn = on_attn - self.to_out = nn.Sequential(nn.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else nn.Linear(inner_dim, dim) + self.to_out = nn.Sequential(ops.Linear(inner_dim, dim * 2), nn.GLU()) if on_attn else ops.Linear(inner_dim, dim) def forward( self, @@ -413,7 +416,7 @@ class AttentionLayers(nn.Module): self.residual_attn = residual_attn self.cross_residual_attn = cross_residual_attn - norm_class = ScaleNorm if use_scalenorm else nn.LayerNorm + norm_class = ScaleNorm if use_scalenorm else ops.LayerNorm norm_class = RMSNorm if use_rmsnorm else norm_class norm_fn = partial(norm_class, dim) @@ -573,13 +576,13 @@ class TransformerWrapper(nn.Module): use_pos_emb and not attn_layers.has_pos_emb) else always(0) self.emb_dropout = nn.Dropout(emb_dropout) - self.project_emb = nn.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity() + self.project_emb = ops.Linear(emb_dim, dim) if emb_dim != dim else nn.Identity() self.attn_layers = attn_layers - self.norm = nn.LayerNorm(dim) + self.norm = ops.LayerNorm(dim) self.init_() - self.to_logits = nn.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t() + self.to_logits = ops.Linear(dim, num_tokens) if not tie_embedding else lambda t: t @ self.token_emb.weight.t() # memory tokens (like [cls]) from Memory Transformers paper num_memory_tokens = default(num_memory_tokens, 0) diff --git a/nodes.py b/nodes.py index a649d84..b44b268 100644 --- a/nodes.py +++ b/nodes.py @@ -557,8 +557,6 @@ class DynamiCrafterI2V: text_emb = positive[0][0].to(device) cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))['last_hidden_state'].to(device) - cond_images = torch.sum(cond_images, dim=0).unsqueeze(0) - cond_images = torch.mean(cond_images, dim=0).unsqueeze(0) img_emb = self.model.image_proj_model(cond_images) @@ -816,12 +814,11 @@ class ToonCrafterInterpolation: pbar = comfy.utils.ProgressBar(len(images) - 1) autocast_condition = (dtype != torch.float32) and not comfy.model_management.is_device_mps(device) with torch.autocast(comfy.model_management.get_autocast_device(device), dtype=dtype) if autocast_condition else nullcontext(): - for i in range(len(images) - 1) if len(images) > 1 else range(len(images)): + for i in range(len(images) - 1): videos, videos2 = None, None mm.soft_empty_cache() image = images[i].unsqueeze(0) - if len(images) !=1: - image2 = images[i+1].unsqueeze(0) + image2 = images[i+1].unsqueeze(0) B, C, H, W = image.shape noise_shape = [B, self.model.model.diffusion_model.out_channels, frames, H // 8, W // 8] @@ -833,16 +830,12 @@ class ToonCrafterInterpolation: image2 += torch.randn_like(image) * augmentation_level encode_pixels = image.unsqueeze(2) * 2 - 1 - videos = encode_pixels # bc1hw - videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames // 2) - - if len(images) == 1: - videos = torch.cat([videos, videos], dim=2) - else: - encode_pixels = image2.unsqueeze(2) * 2 - 1 - videos2 = encode_pixels # bc1hw - videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames // 2) - videos = torch.cat([videos, videos2], dim=2) + videos = encode_pixels # bc1hw + videos = repeat(videos, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) + encode_pixels = image2.unsqueeze(2) * 2 - 1 + videos2 = encode_pixels # bc1hw + videos2 = repeat(videos2, 'b c t h w -> b c (repeat t) h w', repeat=frames//2) + videos = torch.cat([videos, videos2], dim=2) try: z, hs = get_latent_z_with_hidden_states(self.model, videos) @@ -854,25 +847,23 @@ class ToonCrafterInterpolation: img_tensor_repeat = torch.zeros_like(z) img_tensor_repeat[:,:,:1,:,:] = z[:,:,:1,:,:] - if len(images) !=1: - img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:] + img_tensor_repeat[:,:,-1:,:,:] = z[:,:,-1:,:,:] self.model.first_stage_model.to(offload_device) text_emb = positive[0][0].to(device) - self.model.image_proj_model.to(device) cond_images = clip_vision.encode_image(image.permute(0, 2, 3, 1))["last_hidden_state"].to(device) + cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device) + + self.model.image_proj_model.to(device) + img_emb = self.model.image_proj_model(cond_images) - if len(images) !=1: - cond_images2 = clip_vision.encode_image(image2.permute(0, 2, 3, 1))["last_hidden_state"].to(device) - img_emb2 = self.model.image_proj_model(cond_images2) - img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio) - else: - img_embeds = img_emb + img_emb2 = self.model.image_proj_model(cond_images2) + img_embeds = img_emb * image_embed_ratio + img_emb2 * (1.0 - image_embed_ratio) imtext_cond = torch.cat([text_emb, img_embeds], dim=1) - del cond_images, img_emb, text_emb + del cond_images, img_emb, img_emb2, text_emb if comfy.model_management.is_device_mps(device): fs = torch.tensor([fs], dtype=torch.float32, device=self.model.device) diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..640e097 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,15 @@ +[project] +name = "comfyui-dynamicrafterwrapper" +description = "Wrapper nodes to use Dynami/ToonCrafter image2video and frame interpolation models in ComfyUI" +version = "1.0.2" +license = "Apache-2.0" +dependencies = ["einops>=0.3.0", "numpy>=1.24.2", "omegaconf>=2.1.1", "pytorch_lightning>=2.2.1", "tqdm>=4.65.0", "transformers>=4.25.1", "timm"] + +[project.urls] +Repository = "https://github.com/kijai/ComfyUI-DynamiCrafterWrapper" +# Used by Comfy Registry https://comfyregistry.org + +[tool.comfy] +PublisherId = "kijai" +DisplayName = "ComfyUI-DynamiCrafterWrapper" +Icon = "" diff --git a/requirements.txt b/requirements.txt index a4f0627..4111e25 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,10 +1,8 @@ einops>=0.3.0 numpy>=1.24.2 omegaconf>=2.1.1 -Pillow>=9.5.0 pytorch_lightning>=2.2.1 tqdm>=4.65.0 transformers>=4.25.1 timm -open_clip_torch>=2.23.0 -kornia \ No newline at end of file +accelerate \ No newline at end of file diff --git a/scripts/evaluation/funcs.py b/scripts/evaluation/funcs.py index bf61bc5..a1b03b8 100644 --- a/scripts/evaluation/funcs.py +++ b/scripts/evaluation/funcs.py @@ -5,51 +5,39 @@ import torch from einops import rearrange from safetensors.torch import load_file -def load_model_checkpoint(model, ckpt): - def load_checkpoint(model, ckpt, full_strict): - if "safetensors" in ckpt: - try: - state_dict = load_file(ckpt) - except: - state_dict = torch.load(ckpt, map_location="cpu") - else: - state_dict = torch.load(ckpt, map_location="cpu") - if "state_dict" in list(state_dict.keys()): - state_dict = state_dict["state_dict"] +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: + pass - filtered_state_dict = { - k: v - for k, v in state_dict.items() - if not (k.startswith("cond_stage_model") or k.startswith("embedder")) - #if not (k.startswith("cond_stage_model")) - } # Filter out keys starting with "cond_stage_model" and "embedder" +def load_model_checkpoint(model, file_path, dtype, device): + if "safetensors" in file_path: try: - model.load_state_dict(filtered_state_dict, strict=full_strict) + state_dict = load_file(file_path) except: - ## rename the keys for 256x256 model - new_pl_sd = OrderedDict() - for k,v in state_dict.items(): - new_pl_sd[k] = v + state_dict = torch.load(file_path, map_location="cpu") + else: + state_dict = torch.load(file_path, map_location="cpu") + if "state_dict" in list(state_dict.keys()): + state_dict = state_dict["state_dict"] - for k in list(new_pl_sd.keys()): - if "framestride_embed" in k: - new_key = k.replace("framestride_embed", "fps_embedding") - new_pl_sd[new_key] = new_pl_sd[k] - del new_pl_sd[k] - model.load_state_dict(new_pl_sd, strict=full_strict) - # else: - # ## deepspeed - # new_pl_sd = OrderedDict() - # for key in state_dict['module'].keys(): - # new_pl_sd[key[16:]]=state_dict['module'][key] - # model.load_state_dict(new_pl_sd, strict=full_strict) + filtered_state_dict = { + k: v + for k, v in state_dict.items() + if not (k.startswith("cond_stage_model") or k.startswith("embedder")) + #if not (k.startswith("cond_stage_model")) + } # Filter out keys starting with "cond_stage_model" and "embedder" + if is_accelerate_available: + for key in filtered_state_dict: + set_module_tensor_to_device(model, key, dtype=dtype, device=device, value=filtered_state_dict[key]) + else: + model.load_state_dict(filtered_state_dict, strict=True) - return model - load_checkpoint(model, ckpt, full_strict=False) - print('>>> model checkpoint loaded.') return model - def load_prompts(prompt_file): f = open(prompt_file, 'r') prompt_list = []