diff --git a/example_workflows/wanvideo_phantom_subject2vid_example_01.json b/example_workflows/wanvideo_phantom_subject2vid_example_01.json new file mode 100644 index 0000000..96ed432 --- /dev/null +++ b/example_workflows/wanvideo_phantom_subject2vid_example_01.json @@ -0,0 +1,1906 @@ +{ + "id": "c6e410bc-5e2c-460b-ae81-c91b6094fbb1", + "revision": 0, + "last_node_id": 74, + "last_link_id": 120, + "nodes": [ + { + "id": 11, + "type": "LoadWanVideoT5TextEncoder", + "pos": [ + 224.15325927734375, + -34.481563568115234 + ], + "size": [ + 377.1661376953125, + 130 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "wan_t5_model", + "type": "WANTEXTENCODER", + "slot_index": 0, + "links": [ + 15 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "LoadWanVideoT5TextEncoder" + }, + "widgets_values": [ + "umt5-xxl-enc-bf16.safetensors", + "bf16", + "offload_device", + "disabled" + ], + "color": "#332922", + "bgcolor": "#593930" + }, + { + "id": 36, + "type": "Note", + "pos": [ + 723.7317504882812, + -597.3093872070312 + ], + "size": [ + 374.3061828613281, + 171.9547576904297 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "fp8_fast seems to cause huge quality degradation\n\nfp_16_fast enables \"Full FP16 Accmumulation in FP16 GEMMs\" feature available in the very latest pytorch nightly, this is around 20% speed boost. \n\nSageattn if you have it installed can be used for almost double inference speed" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 42, + "type": "Note", + "pos": [ + -165.44613647460938, + -344.9282531738281 + ], + "size": [ + 314.96246337890625, + 152.77333068847656 + ], + "flags": {}, + "order": 2, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "Adjust the blocks to swap based on your VRAM, this is a tradeoff between speed and memory usage.\n\nAlternatively there's option to use VRAM management introduced in DiffSynt-Studios. This is usually slower, but saves even more VRAM compared to BlockSwap" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 50, + "type": "CLIPTextEncode", + "pos": [ + -78.64810180664062, + 1769.301513671875 + ], + "size": [ + 400, + 200 + ], + "flags": {}, + "order": 18, + "mode": 2, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 53 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "slot_index": 0, + "links": [ + 55 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.29", + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 48, + "type": "CLIPLoader", + "pos": [ + -438.6482238769531, + 1519.30126953125 + ], + "size": [ + 315, + 106 + ], + "flags": {}, + "order": 3, + "mode": 2, + "inputs": [], + "outputs": [ + { + "name": "CLIP", + "type": "CLIP", + "slot_index": 0, + "links": [ + 52, + 53 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.29", + "Node name for S&R": "CLIPLoader" + }, + "widgets_values": [ + "umt5_xxl_fp16.safetensors", + "wan", + "default" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 51, + "type": "Note", + "pos": [ + -408.648193359375, + 1349.3011474609375 + ], + "size": [ + 253.16725158691406, + 88 + ], + "flags": {}, + "order": 4, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "You can also use native ComfyUI text encoding with these nodes instead of the original, the models are node specific and can't otherwise be mixed." + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 49, + "type": "CLIPTextEncode", + "pos": [ + -78.64810180664062, + 1519.30126953125 + ], + "size": [ + 400, + 200 + ], + "flags": {}, + "order": 17, + "mode": 2, + "inputs": [ + { + "name": "clip", + "type": "CLIP", + "link": 52 + } + ], + "outputs": [ + { + "name": "CONDITIONING", + "type": "CONDITIONING", + "slot_index": 0, + "links": [ + 54 + ] + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.29", + "Node name for S&R": "CLIPTextEncode" + }, + "widgets_values": [ + "high quality nature video featuring a red panda balancing on a bamboo stem while a bird lands on it's head, on the background there is a waterfall" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 45, + "type": "WanVideoVRAMManagement", + "pos": [ + -158.19737243652344, + -136.97467041015625 + ], + "size": [ + 315, + 58 + ], + "flags": {}, + "order": 5, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vram_management_args", + "type": "VRAM_MANAGEMENTARGS", + "links": [] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoVRAMManagement" + }, + "widgets_values": [ + 1 + ] + }, + { + "id": 33, + "type": "Note", + "pos": [ + -153.7365264892578, + -16.124788284301758 + ], + "size": [ + 359.0753479003906, + 88 + ], + "flags": {}, + "order": 6, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 44, + "type": "Note", + "pos": [ + -98.58364868164062, + -675.3411254882812 + ], + "size": [ + 303.0501403808594, + 88 + ], + "flags": {}, + "order": 7, + "mode": 0, + "inputs": [], + "outputs": [], + "properties": {}, + "widgets_values": [ + "If you have Triton installed, connect this for ~30% speed increase" + ], + "color": "#432", + "bgcolor": "#653" + }, + { + "id": 39, + "type": "WanVideoBlockSwap", + "pos": [ + 253.16395568847656, + -343.3807678222656 + ], + "size": [ + 315, + 154 + ], + "flags": {}, + "order": 8, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "block_swap_args", + "type": "BLOCKSWAPARGS", + "slot_index": 0, + "links": [] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoBlockSwap" + }, + "widgets_values": [ + 20, + false, + false, + true, + 0 + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 38, + "type": "WanVideoVAELoader", + "pos": [ + 1687.4093017578125, + -582.2750854492188 + ], + "size": [ + 315, + 82 + ], + "flags": {}, + "order": 9, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "vae", + "type": "WANVAE", + "slot_index": 0, + "links": [ + 43, + 59, + 110 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoVAELoader" + }, + "widgets_values": [ + "wanvideo\\Wan2_1_VAE_bf16.safetensors", + "bf16" + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 46, + "type": "WanVideoTextEmbedBridge", + "pos": [ + 371.3523254394531, + 1509.30126953125 + ], + "size": [ + 315, + 46 + ], + "flags": {}, + "order": 22, + "mode": 2, + "inputs": [ + { + "name": "positive", + "type": "CONDITIONING", + "link": 54 + }, + { + "name": "negative", + "type": "CONDITIONING", + "link": 55 + } + ], + "outputs": [ + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "links": null + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoTextEmbedBridge" + }, + "widgets_values": [] + }, + { + "id": 28, + "type": "WanVideoDecode", + "pos": [ + 1692.973876953125, + -404.8614501953125 + ], + "size": [ + 315, + 174 + ], + "flags": {}, + "order": 30, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 43 + }, + { + "name": "samples", + "type": "LATENT", + "link": 33 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "slot_index": 0, + "links": [ + 81 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoDecode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128 + ], + "color": "#322", + "bgcolor": "#533" + }, + { + "id": 35, + "type": "WanVideoTorchCompileSettings", + "pos": [ + 222.5817413330078, + -677.6240844726562 + ], + "size": [ + 421.6000061035156, + 202 + ], + "flags": {}, + "order": 10, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "torch_compile_args", + "type": "WANCOMPILEARGS", + "slot_index": 0, + "links": [ + 70 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoTorchCompileSettings" + }, + "widgets_values": [ + "inductor", + false, + "default", + false, + 64, + true, + 128 + ] + }, + { + "id": 64, + "type": "WanVideoTeaCache", + "pos": [ + 1203.9754638671875, + -657.2056884765625 + ], + "size": [ + 380.4000244140625, + 178 + ], + "flags": {}, + "order": 11, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "teacache_args", + "type": "TEACACHEARGS", + "links": [ + 71 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a623f87dcad9cff5a690559fe559566be4045a9a", + "Node name for S&R": "WanVideoTeaCache" + }, + "widgets_values": [ + 0.10000000000000002, + 6, + -1, + "offload_device", + true, + "e0" + ] + }, + { + "id": 22, + "type": "WanVideoModelLoader", + "pos": [ + 620.3950805664062, + -357.8426818847656 + ], + "size": [ + 477.4410095214844, + 234 + ], + "flags": {}, + "order": 19, + "mode": 0, + "inputs": [ + { + "name": "compile_args", + "shape": 7, + "type": "WANCOMPILEARGS", + "link": 70 + }, + { + "name": "block_swap_args", + "shape": 7, + "type": "BLOCKSWAPARGS", + "link": null + }, + { + "name": "lora", + "shape": 7, + "type": "WANVIDLORA", + "link": null + }, + { + "name": "vram_management_args", + "shape": 7, + "type": "VRAM_MANAGEMENTARGS", + "link": null + }, + { + "name": "vace_model", + "shape": 7, + "type": "VACEPATH", + "link": null + } + ], + "outputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "slot_index": 0, + "links": [ + 29 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoModelLoader" + }, + "widgets_values": [ + "WanVideo\\Phantom-Wan-1_3B_fp16.safetensors", + "fp16_fast", + "disabled", + "main_device", + "sageattn" + ], + "color": "#223", + "bgcolor": "#335" + }, + { + "id": 61, + "type": "INTConstant", + "pos": [ + -507.47119140625, + -30.832414627075195 + ], + "size": [ + 210, + 58 + ], + "flags": {}, + "order": 12, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "value", + "type": "INT", + "links": [ + 76, + 79, + 85, + 89 + ] + } + ], + "title": "Width", + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f", + "Node name for S&R": "INTConstant" + }, + "widgets_values": [ + 1280 + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 62, + "type": "INTConstant", + "pos": [ + -501.34527587890625, + 83.66830444335938 + ], + "size": [ + 210, + 58 + ], + "flags": {}, + "order": 13, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "value", + "type": "INT", + "links": [ + 77, + 80, + 86, + 90 + ] + } + ], + "title": "Height", + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f", + "Node name for S&R": "INTConstant" + }, + "widgets_values": [ + 768 + ], + "color": "#1b4669", + "bgcolor": "#29699c" + }, + { + "id": 66, + "type": "ImageConcatMulti", + "pos": [ + 1772.1153564453125, + -17.343990325927734 + ], + "size": [ + 315, + 150 + ], + "flags": {}, + "order": 31, + "mode": 0, + "inputs": [ + { + "name": "image_1", + "type": "IMAGE", + "link": 81 + }, + { + "name": "image_2", + "type": "IMAGE", + "link": 82 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 83 + ] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f" + }, + "widgets_values": [ + 2, + "up", + false, + null + ] + }, + { + "id": 67, + "type": "LoadImage", + "pos": [ + -529.6270751953125, + 642.4041137695312 + ], + "size": [ + 315, + 314 + ], + "flags": {}, + "order": 14, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 87 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.29", + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "oldman_upscaled.png", + "image" + ] + }, + { + "id": 68, + "type": "ImageResizeKJ", + "pos": [ + -136.8474884033203, + 636.0980224609375 + ], + "size": [ + 315, + 238 + ], + "flags": {}, + "order": 20, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 87 + }, + { + "name": "width_input", + "shape": 7, + "type": "INT", + "link": null + }, + { + "name": "height_input", + "shape": 7, + "type": "INT", + "link": null + }, + { + "name": "get_image_size", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 85 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 86 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 91 + ] + }, + { + "name": "width", + "type": "INT", + "links": [] + }, + { + "name": "height", + "type": "INT", + "links": [] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f", + "Node name for S&R": "ImageResizeKJ" + }, + "widgets_values": [ + 512, + 512, + "lanczos", + true, + 8, + "disabled" + ] + }, + { + "id": 65, + "type": "ImageResizeKJ", + "pos": [ + -138.123046875, + 268.34466552734375 + ], + "size": [ + 315, + 238 + ], + "flags": {}, + "order": 21, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 72 + }, + { + "name": "width_input", + "shape": 7, + "type": "INT", + "link": null + }, + { + "name": "height_input", + "shape": 7, + "type": "INT", + "link": null + }, + { + "name": "get_image_size", + "shape": 7, + "type": "IMAGE", + "link": null + }, + { + "name": "width", + "type": "INT", + "widget": { + "name": "width" + }, + "link": 79 + }, + { + "name": "height", + "type": "INT", + "widget": { + "name": "height" + }, + "link": 80 + } + ], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 73 + ] + }, + { + "name": "width", + "type": "INT", + "links": [] + }, + { + "name": "height", + "type": "INT", + "links": [] + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f", + "Node name for S&R": "ImageResizeKJ" + }, + "widgets_values": [ + 512, + 512, + "lanczos", + true, + 8, + "disabled" + ] + }, + { + "id": 57, + "type": "LoadImage", + "pos": [ + -549.6646728515625, + 251.52166748046875 + ], + "size": [ + 315, + 314 + ], + "flags": {}, + "order": 15, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "IMAGE", + "type": "IMAGE", + "links": [ + 72 + ] + }, + { + "name": "MASK", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.29", + "Node name for S&R": "LoadImage" + }, + "widgets_values": [ + "anya.webp", + "image" + ] + }, + { + "id": 16, + "type": "WanVideoTextEncode", + "pos": [ + 675.8850708007812, + -36.032100677490234 + ], + "size": [ + 420.30511474609375, + 261.5306701660156 + ], + "flags": {}, + "order": 16, + "mode": 0, + "inputs": [ + { + "name": "t5", + "type": "WANTEXTENCODER", + "link": 15 + }, + { + "name": "model_to_offload", + "shape": 7, + "type": "WANVIDEOMODEL", + "link": null + } + ], + "outputs": [ + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "slot_index": 0, + "links": [ + 30 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoTextEncode" + }, + "widgets_values": [ + "an old man is playing with a chibi anime figurine", + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", + true + ], + "color": "#332922", + "bgcolor": "#593930" + }, + { + "id": 63, + "type": "PreviewImage", + "pos": [ + 1249.652587890625, + 640.6929931640625 + ], + "size": [ + 675.1277465820312, + 258 + ], + "flags": {}, + "order": 25, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 102 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.29", + "Node name for S&R": "PreviewImage" + }, + "widgets_values": [] + }, + { + "id": 69, + "type": "ImagePadKJ", + "pos": [ + 241.51797485351562, + 614.47705078125 + ], + "size": [ + 315, + 262 + ], + "flags": {}, + "order": 23, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 91 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + }, + { + "name": "target_width", + "shape": 7, + "type": "INT", + "link": 89 + }, + { + "name": "target_height", + "shape": 7, + "type": "INT", + "link": 90 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 102, + 105 + ] + }, + { + "name": "masks", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f", + "Node name for S&R": "ImagePadKJ" + }, + "widgets_values": [ + 0, + 0, + 0, + 0, + 0, + "color", + "255,255,255" + ] + }, + { + "id": 60, + "type": "ImagePadKJ", + "pos": [ + 241.20155334472656, + 261.9258728027344 + ], + "size": [ + 315, + 262 + ], + "flags": {}, + "order": 24, + "mode": 0, + "inputs": [ + { + "name": "image", + "type": "IMAGE", + "link": 73 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + }, + { + "name": "target_width", + "shape": 7, + "type": "INT", + "link": 76 + }, + { + "name": "target_height", + "shape": 7, + "type": "INT", + "link": 77 + } + ], + "outputs": [ + { + "name": "images", + "type": "IMAGE", + "links": [ + 82, + 106 + ] + }, + { + "name": "masks", + "type": "MASK", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-kjnodes", + "ver": "3e3a1a8aac61dc4515f6a7da74e026f05a80299f", + "Node name for S&R": "ImagePadKJ" + }, + "widgets_values": [ + 0, + 0, + 0, + 0, + 0, + "color", + "255,255,255" + ] + }, + { + "id": 74, + "type": "WanVideoPhantomEmbeds", + "pos": [ + 1243.949462890625, + 424.87481689453125 + ], + "size": [ + 380.4000244140625, + 142 + ], + "flags": {}, + "order": 28, + "mode": 0, + "inputs": [ + { + "name": "phantom_latent_1", + "type": "LATENT", + "link": 119 + }, + { + "name": "phantom_latent_2", + "shape": 7, + "type": "LATENT", + "link": 120 + }, + { + "name": "phantom_latent_3", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "phantom_latent_4", + "shape": 7, + "type": "LATENT", + "link": null + } + ], + "outputs": [ + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "links": [ + 114 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "1f535743870da83c530386874f10aabc70201919", + "Node name for S&R": "WanVideoPhantomEmbeds" + }, + "widgets_values": [ + 81, + 5 + ] + }, + { + "id": 73, + "type": "WanVideoEncode", + "pos": [ + 685.5361938476562, + 634.3781127929688 + ], + "size": [ + 330, + 242 + ], + "flags": {}, + "order": 26, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 110 + }, + { + "name": "image", + "type": "IMAGE", + "link": 105 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 119 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a623f87dcad9cff5a690559fe559566be4045a9a", + "Node name for S&R": "WanVideoEncode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + 0, + 1 + ] + }, + { + "id": 56, + "type": "WanVideoEncode", + "pos": [ + 688.443359375, + 299.2251892089844 + ], + "size": [ + 330, + 242 + ], + "flags": {}, + "order": 27, + "mode": 0, + "inputs": [ + { + "name": "vae", + "type": "WANVAE", + "link": 59 + }, + { + "name": "image", + "type": "IMAGE", + "link": 106 + }, + { + "name": "mask", + "shape": 7, + "type": "MASK", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 120 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "a623f87dcad9cff5a690559fe559566be4045a9a", + "Node name for S&R": "WanVideoEncode" + }, + "widgets_values": [ + false, + 272, + 272, + 144, + 128, + 0, + 1 + ] + }, + { + "id": 27, + "type": "WanVideoSampler", + "pos": [ + 1315.2401123046875, + -401.48028564453125 + ], + "size": [ + 315, + 729 + ], + "flags": {}, + "order": 29, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 29 + }, + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "link": 30 + }, + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 114 + }, + { + "name": "samples", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "feta_args", + "shape": 7, + "type": "FETAARGS", + "link": null + }, + { + "name": "context_options", + "shape": 7, + "type": "WANVIDCONTEXT", + "link": null + }, + { + "name": "teacache_args", + "shape": 7, + "type": "TEACACHEARGS", + "link": 71 + }, + { + "name": "flowedit_args", + "shape": 7, + "type": "FLOWEDITARGS", + "link": null + }, + { + "name": "slg_args", + "shape": 7, + "type": "SLGARGS", + "link": null + }, + { + "name": "loop_args", + "shape": 7, + "type": "LOOPARGS", + "link": null + }, + { + "name": "experimental_args", + "shape": 7, + "type": "EXPERIMENTALARGS", + "link": null + }, + { + "name": "sigmas", + "shape": 7, + "type": "SIGMAS", + "link": null + }, + { + "name": "unianimate_poses", + "shape": 7, + "type": "UNIANIMATE_POSE", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "slot_index": 0, + "links": [ + 33 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "6099ad393b071728032fd481e96d77d2900eee2c", + "Node name for S&R": "WanVideoSampler" + }, + "widgets_values": [ + 40, + 7.500000000000002, + 5, + 44, + "fixed", + true, + "unipc", + 0, + 1, + false, + "comfy", + "" + ] + }, + { + "id": 30, + "type": "VHS_VideoCombine", + "pos": [ + 2202.927001953125, + -570.5418701171875 + ], + "size": [ + 1245.8460693359375, + 1819.0152587890625 + ], + "flags": {}, + "order": 32, + "mode": 0, + "inputs": [ + { + "name": "images", + "type": "IMAGE", + "link": 83 + }, + { + "name": "audio", + "shape": 7, + "type": "AUDIO", + "link": null + }, + { + "name": "meta_batch", + "shape": 7, + "type": "VHS_BatchManager", + "link": null + }, + { + "name": "vae", + "shape": 7, + "type": "VAE", + "link": null + } + ], + "outputs": [ + { + "name": "Filenames", + "type": "VHS_FILENAMES", + "links": null + } + ], + "properties": { + "cnr_id": "comfyui-videohelpersuite", + "ver": "0a75c7958fe320efcb052f1d9f8451fd20c730a8", + "Node name for S&R": "VHS_VideoCombine" + }, + "widgets_values": { + "frame_rate": 16, + "loop_count": 0, + "filename_prefix": "WanVideo21_Phantom", + "format": "video/h264-mp4", + "pix_fmt": "yuv420p", + "crf": 19, + "save_metadata": true, + "trim_to_audio": false, + "pingpong": false, + "save_output": false, + "videopreview": { + "hidden": false, + "paused": false, + "params": { + "filename": "WanVideo2_1_T2V_00013.mp4", + "subfolder": "", + "type": "temp", + "format": "video/h264-mp4", + "frame_rate": 16, + "workflow": "WanVideo2_1_T2V_00013.png", + "fullpath": "N:\\AI\\ComfyUI\\temp\\WanVideo2_1_T2V_00013.mp4" + } + } + } + } + ], + "links": [ + [ + 15, + 11, + 0, + 16, + 0, + "WANTEXTENCODER" + ], + [ + 29, + 22, + 0, + 27, + 0, + "WANVIDEOMODEL" + ], + [ + 30, + 16, + 0, + 27, + 1, + "WANVIDEOTEXTEMBEDS" + ], + [ + 33, + 27, + 0, + 28, + 1, + "LATENT" + ], + [ + 43, + 38, + 0, + 28, + 0, + "VAE" + ], + [ + 52, + 48, + 0, + 49, + 0, + "CLIP" + ], + [ + 53, + 48, + 0, + 50, + 0, + "CLIP" + ], + [ + 54, + 49, + 0, + 46, + 0, + "CONDITIONING" + ], + [ + 55, + 50, + 0, + 46, + 1, + "CONDITIONING" + ], + [ + 59, + 38, + 0, + 56, + 0, + "WANVAE" + ], + [ + 70, + 35, + 0, + 22, + 0, + "WANCOMPILEARGS" + ], + [ + 71, + 64, + 0, + 27, + 6, + "TEACACHEARGS" + ], + [ + 72, + 57, + 0, + 65, + 0, + "IMAGE" + ], + [ + 73, + 65, + 0, + 60, + 0, + "IMAGE" + ], + [ + 76, + 61, + 0, + 60, + 2, + "INT" + ], + [ + 77, + 62, + 0, + 60, + 3, + "INT" + ], + [ + 79, + 61, + 0, + 65, + 4, + "INT" + ], + [ + 80, + 62, + 0, + 65, + 5, + "INT" + ], + [ + 81, + 28, + 0, + 66, + 0, + "IMAGE" + ], + [ + 82, + 60, + 0, + 66, + 1, + "IMAGE" + ], + [ + 83, + 66, + 0, + 30, + 0, + "IMAGE" + ], + [ + 85, + 61, + 0, + 68, + 4, + "INT" + ], + [ + 86, + 62, + 0, + 68, + 5, + "INT" + ], + [ + 87, + 67, + 0, + 68, + 0, + "IMAGE" + ], + [ + 89, + 61, + 0, + 69, + 2, + "INT" + ], + [ + 90, + 62, + 0, + 69, + 3, + "INT" + ], + [ + 91, + 68, + 0, + 69, + 0, + "IMAGE" + ], + [ + 102, + 69, + 0, + 63, + 0, + "IMAGE" + ], + [ + 105, + 69, + 0, + 73, + 1, + "IMAGE" + ], + [ + 106, + 60, + 0, + 56, + 1, + "IMAGE" + ], + [ + 110, + 38, + 0, + 73, + 0, + "WANVAE" + ], + [ + 114, + 74, + 0, + 27, + 2, + "WANVIDIMAGE_EMBEDS" + ], + [ + 119, + 73, + 0, + 74, + 0, + "LATENT" + ], + [ + 120, + 56, + 0, + 74, + 1, + "LATENT" + ] + ], + "groups": [ + { + "id": 1, + "title": "ComfyUI text encoding alternative", + "bounding": [ + -501.4642639160156, + 1205.3677978515625, + 1210.621337890625, + 805.9080810546875 + ], + "color": "#3f789e", + "font_size": 24, + "flags": {} + } + ], + "config": {}, + "extra": { + "ds": { + "scale": 0.6115909044841836, + "offset": [ + 80.2022437434934, + 871.50326287299 + ] + }, + "frontendVersion": "1.17.3", + "node_versions": { + "ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd", + "comfy-core": "0.3.26", + "ComfyUI-VideoHelperSuite": "0a75c7958fe320efcb052f1d9f8451fd20c730a8" + }, + "VHS_latentpreview": true, + "VHS_latentpreviewrate": 0, + "VHS_MetadataImage": true, + "VHS_KeepIntermediate": true + }, + "version": 0.4 +} \ No newline at end of file diff --git a/fp8_optimization.py b/fp8_optimization.py index 0688ee6..f32eae3 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -7,8 +7,8 @@ def fp8_linear_forward(cls, original_dtype, input): weight_dtype = cls.weight.dtype if weight_dtype in [torch.float8_e4m3fn, torch.float8_e5m2]: if len(input.shape) == 3: - target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn - inn = input.reshape(-1, input.shape[2]).to(target_dtype) + #target_dtype = torch.float8_e5m2 if weight_dtype == torch.float8_e4m3fn else torch.float8_e4m3fn + inn = input.reshape(-1, input.shape[2]).to(weight_dtype) w = cls.weight.t() scale = torch.ones((1), device=input.device, dtype=torch.float32) diff --git a/nodes.py b/nodes.py index 47de9e9..442f972 100644 --- a/nodes.py +++ b/nodes.py @@ -465,7 +465,7 @@ class WanVideoModelLoader: "model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' -folder",}), "base_precision": (["fp32", "bf16", "fp16", "fp16_fast"], {"default": "bf16"}), - "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}), + "quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn', 'torchao_fp8dq', "torchao_fp8dqrow", "torchao_int8dq", "torchao_fp6", "torchao_int4", "torchao_int8"], {"default": 'disabled', "tooltip": "optional quantization method"}), "load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}), }, "optional": { @@ -621,6 +621,8 @@ class WanVideoModelLoader: "vace_layers": vace_layers, "vace_in_dim": vace_in_dim, "inject_sample_info": True if "fps_embedding.weight" in sd else False, + "add_ref_conv": True if "ref_conv.weight" in sd else False, + "in_dim_ref_conv": sd["ref_conv.weight"].shape[1] if "ref_conv.weight" in sd else None, } with init_empty_weights(): @@ -637,7 +639,7 @@ class WanVideoModelLoader: block.cam_encoder.bias.data.zero_() block.projector.weight = nn.Parameter(torch.eye(dim)) block.projector.bias = nn.Parameter(torch.zeros(dim)) - + comfy_model = WanVideoModel( WanVideoModelConfig(base_dtype), model_type=comfy.model_base.ModelType.FLOW, @@ -646,13 +648,13 @@ class WanVideoModelLoader: if not "torchao" in quantization: - if quantization == "fp8_e4m3fn" or quantization == "fp8_e4m3fn_fast" or quantization == "fp8_scaled": + if "fp8_e4m3fn" in quantization: dtype = torch.float8_e4m3fn elif quantization == "fp8_e5m2": dtype = torch.float8_e5m2 else: dtype = base_dtype - params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation"} + params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding"} #if lora is not None: # transformer_load_device = device if not lora_low_mem_load: @@ -663,7 +665,7 @@ class WanVideoModelLoader: total=param_count, leave=True): dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype - if "modulation" in name or "time_" in name: + if "patch_embedding" in name: dtype_to_use = torch.float32 set_module_tensor_to_device(transformer, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name]) comfy_model.diffusion_model = transformer @@ -708,7 +710,7 @@ class WanVideoModelLoader: transformer.patch_embedding.kernel_size, transformer.patch_embedding.stride, transformer.patch_embedding.padding, - ).to(device=device, dtype=torch.bfloat16) + ).to(device=device, dtype=torch.float32) new_in.weight.zero_() new_in.bias.zero_() @@ -728,13 +730,16 @@ class WanVideoModelLoader: #patcher.load(device, full_load=True) patcher.model.is_patched = True - del sd - if quantization == "fp8_e4m3fn_fast": + + if "fast" in quantization: from .fp8_optimization import convert_fp8_linear - #params_to_keep.update({"ffn"}) + if quantization == "fp8_e4m3fn_fast_no_ffn": + params_to_keep.update({"ffn"}) print(params_to_keep) - convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep) + convert_fp8_linear(patcher.model.diffusion_model, base_dtype, params_to_keep=params_to_keep, sd=sd) + + del sd if vram_management_args is not None: from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear @@ -864,7 +869,7 @@ class WanVideoModelLoader: patcher.model["base_path"] = model_path patcher.model["model_name"] = model patcher.model["manual_offloading"] = manual_offloading - patcher.model["quantization"] = "disabled" + patcher.model["quantization"] = quantization patcher.model["auto_cpu_offload"] = True if vram_management_args is not None else False patcher.model["control_lora"] = control_lora @@ -1750,6 +1755,68 @@ class WanVideoEmptyEmbeds: return (embeds,) +# region phantom +class WanVideoPhantomEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}), + "phantom_latent_1": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + + "phantom_cfg_scale": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "CFG scale for the extra phantom cond pass"}), + "phantom_start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the phantom model"}), + "phantom_end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the phantom model"}), + }, + "optional": { + "phantom_latent_2": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_latent_3": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "phantom_latent_4": ("LATENT", {"tooltip": "reference latents for the phantom model"}), + "vace_embeds": ("WANVIDIMAGE_EMBEDS", {"tooltip": "VACE embeds"}), + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", ) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "process" + CATEGORY = "WanVideoWrapper" + + def process(self, num_frames, phantom_cfg_scale, phantom_start_percent, phantom_end_percent, phantom_latent_1, phantom_latent_2=None, phantom_latent_3=None, phantom_latent_4=None, vace_embeds=None): + vae_stride = (4, 8, 8) + samples = phantom_latent_1["samples"].squeeze(0) + if phantom_latent_2 is not None: + samples = torch.cat([samples, phantom_latent_2["samples"].squeeze(0)], dim=1) + if phantom_latent_3 is not None: + samples = torch.cat([samples, phantom_latent_3["samples"].squeeze(0)], dim=1) + if phantom_latent_4 is not None: + samples = torch.cat([samples, phantom_latent_4["samples"].squeeze(0)], dim=1) + C, T, H, W = samples.shape + + target_shape = (16, (num_frames - 1) // vae_stride[0] + 1 + T, + H * 8 // vae_stride[1], + W * 8 // vae_stride[2]) + + embeds = { + "target_shape": target_shape, + "num_frames": num_frames, + "phantom_latents": samples, + "phantom_cfg_scale": phantom_cfg_scale, + "phantom_start_percent": phantom_start_percent, + "phantom_end_percent": phantom_end_percent, + } + if vace_embeds is not None: + vace_input = { + "vace_context": vace_embeds["vace_context"], + "vace_scale": vace_embeds["vace_scale"], + "has_ref": vace_embeds["has_ref"], + "vace_start_percent": vace_embeds["vace_start_percent"], + "vace_end_percent": vace_embeds["vace_end_percent"], + "vace_seq_len": vace_embeds["vace_seq_len"], + "additional_vace_inputs": vace_embeds["additional_vace_inputs"], + } + embeds.update(vace_input) + + return (embeds,) + class WanVideoControlEmbeds: @classmethod def INPUT_TYPES(s): @@ -1758,6 +1825,9 @@ class WanVideoControlEmbeds: "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the control signal"}), "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the control signal"}), }, + "optional": { + "fun_ref_image": ("LATENT", {"tooltip": "Reference latent for the Fun 1.1 -model"}), + } } RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", ) @@ -1765,7 +1835,7 @@ class WanVideoControlEmbeds: FUNCTION = "process" CATEGORY = "WanVideoWrapper" - def process(self, latents, start_percent, end_percent): + def process(self, latents, start_percent, end_percent, fun_ref_image=None): samples = latents["samples"].squeeze(0) C, T, H, W = samples.shape @@ -1780,7 +1850,8 @@ class WanVideoControlEmbeds: "control_embeds": { "control_images": samples, "start_percent": start_percent, - "end_percent": end_percent + "end_percent": end_percent, + "fun_ref_image": fun_ref_image["samples"][:,:, 0] if fun_ref_image is not None else None, } } @@ -2213,7 +2284,7 @@ class WanVideoSampler: patcher = model model = model.model transformer = model.diffusion_model - + dtype = model["dtype"] control_lora = model["control_lora"] device = mm.get_torch_device() @@ -2269,7 +2340,9 @@ class WanVideoSampler: control_latents, clip_fea, clip_fea_neg, end_image, recammaster, camera_embed, unianim_data = None, None, None, None, None, None, None vace_data, vace_context, vace_scale = None, None, None - fun_or_fl2v_model, has_ref, drop_last = False, False, False + fun_or_fl2v_model, has_ref, drop_last, = False, False, False + phantom_latents = None + fun_ref_image = None image_cond = image_embeds.get("image_embeds", None) @@ -2292,7 +2365,11 @@ class WanVideoSampler: image_cond = image_embeds.get("image_embeds", None) print("image_cond", image_cond.shape) clip_fea = image_embeds.get("clip_context", None) + if clip_fea is not None: + clip_fea = clip_fea.to(dtype) clip_fea_neg = image_embeds.get("negative_clip_context", None) + if clip_fea_neg is not None: + clip_fea_neg = clip_fea_neg.to(dtype) control_embeds = image_embeds.get("control_embeds", None) if control_embeds is not None: @@ -2371,7 +2448,7 @@ class WanVideoSampler: raise ValueError("Control signal only works with Fun-Control model") image_cond = torch.zeros_like(control_latents).to(device) #fun control clip_fea = None - + fun_ref_image = control_embeds.get("fun_ref_image", None) control_start_percent = control_embeds.get("start_percent", 0.0) control_end_percent = control_embeds.get("end_percent", 1.0) else: @@ -2382,6 +2459,13 @@ class WanVideoSampler: masked_video_latents_input = torch.zeros_like(noise) image_cond = torch.cat([mask_latents, masked_video_latents_input], dim=0).to(device) + phantom_latents = image_embeds.get("phantom_latents", None) + phantom_cfg_scale = image_embeds.get("phantom_cfg_scale", None) + phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0) + phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0) + if phantom_latents is not None: + phantom_latents = phantom_latents.to(device) + latent_video_length = noise.shape[1] if unianimate_poses is not None: @@ -2577,6 +2661,7 @@ class WanVideoSampler: transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] transformer.teacache_start_step = teacache_args["start_step"] transformer.teacache_cache_device = teacache_args["cache_device"] + log.info(f"TeaCache: Using cache device: {transformer.teacache_state.cache_device}") transformer.teacache_end_step = len(timesteps)-1 if teacache_args["end_step"] == -1 else teacache_args["end_step"] transformer.teacache_use_coefficients = teacache_args["use_coefficients"] transformer.teacache_mode = teacache_args["mode"] @@ -2593,6 +2678,9 @@ class WanVideoSampler: transformer.slg_blocks = None self.teacache_state = [None, None] + if phantom_latents is not None: + log.info(f"Phantom latents shape: {phantom_latents.shape}") + self.teacache_state = [None, None, None] self.teacache_state_source = [None, None] self.teacache_states_context = [] @@ -2601,6 +2689,8 @@ class WanVideoSampler: source_image_embeds = flowedit_args.get("source_image_embeds", image_embeds) source_image_cond = source_image_embeds.get("image_embeds", None) source_clip_fea = source_image_embeds.get("clip_fea", clip_fea) + if source_image_cond is not None: + source_image_cond = source_image_cond.to(dtype) skip_steps = flowedit_args["skip_steps"] drift_steps = flowedit_args["drift_steps"] source_cfg = flowedit_args["source_cfg"] @@ -2649,7 +2739,8 @@ class WanVideoSampler: #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, control_latents=None, vace_data=None, unianim_data=None, teacache_state=None): - with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): + z = z.to(dtype) + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: return latent_model_input*0, None @@ -2664,9 +2755,14 @@ class WanVideoSampler: else: if (control_start_percent <= current_step_percentage <= control_end_percent) or \ (control_end_percent > 0 and idx == 0 and current_step_percentage >= control_start_percent): - image_cond_input = torch.cat([control_latents, image_cond]) + image_cond_input = torch.cat([control_latents.to(z), image_cond.to(z)]) else: - image_cond_input = torch.cat([torch.zeros_like(image_cond), image_cond]) + image_cond_input = torch.cat([torch.zeros_like(image_cond, dtype=dtype), image_cond.to(z)]) + if fun_ref_image is not None: + fun_ref_input = fun_ref_image.to(z) + else: + fun_ref_input = torch.zeros_like(z, dtype=z.dtype)[:, 0].unsqueeze(1) + #fun_ref_input = None if control_lora: if not control_start_percent <= current_step_percentage <= control_end_percent: @@ -2676,17 +2772,30 @@ class WanVideoSampler: patcher.unpatch_model(device) patcher.model.is_patched = False else: - image_cond_input = control_latents.to(device) + image_cond_input = control_latents.to(z) if not patcher.model.is_patched: log.info("Loading LoRA...") patcher = apply_lora(patcher, device, device, low_mem_load=False) patcher.model.is_patched = True else: - image_cond_input = image_cond + image_cond_input = image_cond.to(z) if image_cond is not None else None if recammaster is not None: z = torch.cat([z, recam_latents.to(z)], dim=1) - + use_phantom = False + if phantom_latents is not None: + if (phantom_start_percent <= current_step_percentage <= phantom_end_percent) or \ + (phantom_end_percent > 0 and idx == 0 and current_step_percentage >= phantom_start_percent): + + z_pos = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) + z_phantom_img = torch.cat([z[:,:-phantom_latents.shape[1]], phantom_latents.to(z)], dim=1) + z_neg = torch.cat([z[:,:-phantom_latents.shape[1]], torch.zeros_like(phantom_latents).to(z)], dim=1) + use_phantom = True + if len(teacache_state) != 3: + teacache_state.append(None) + if not use_phantom: + z_pos = z_neg = z + base_params = { 'seq_len': seq_len, 'device': device, @@ -2694,9 +2803,9 @@ class WanVideoSampler: 't': timestep, 'current_step': idx, 'control_lora_enabled': control_lora_enabled, - 'vace_data': vace_data, 'camera_embed': camera_embed, 'unianim_data': unianim_data, + 'fun_ref': fun_ref_input if fun_ref_image is not None else None, } batch_size = 1 @@ -2707,9 +2816,10 @@ class WanVideoSampler: if not batched_cfg: #cond noise_pred_cond, teacache_state_cond = transformer( - [z], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, + [z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None, clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage, pred_id=teacache_state[0] if teacache_state else None, + vace_data=vace_data, **base_params ) noise_pred_cond = noise_pred_cond[0].to(intermediate_device) @@ -2724,13 +2834,28 @@ class WanVideoSampler: return noise_pred_cond, [teacache_state_cond] #uncond noise_pred_uncond, teacache_state_uncond = transformer( - [z], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + [z_neg], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, y=[image_cond_input] if image_cond_input is not None else None, is_uncond=True, current_step_percentage=current_step_percentage, pred_id=teacache_state[1] if teacache_state else None, + vace_data=vace_data, **base_params ) noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device) + #phantom + if use_phantom: + noise_pred_phantom, teacache_state_phantom = transformer( + [z_phantom_img], context=negative_embeds, clip_fea=clip_fea_neg if clip_fea_neg is not None else clip_fea, + y=[image_cond_input] if image_cond_input is not None else None, + is_uncond=True, current_step_percentage=current_step_percentage, + pred_id=teacache_state[2] if teacache_state else None, + vace_data=None, + **base_params + ) + noise_pred_phantom = noise_pred_phantom[0].to(intermediate_device) + + noise_pred = noise_pred_uncond + phantom_cfg_scale * (noise_pred_phantom - noise_pred_uncond) + cfg_scale * (noise_pred_cond - noise_pred_phantom) + return noise_pred, [teacache_state_cond, teacache_state_uncond, teacache_state_phantom] #batched else: teacache_state_uncond = None @@ -2827,11 +2952,11 @@ class WanVideoSampler: latent_model_input = torch.cat([latent_model_input[:, shift_idx:]] + [latent_model_input[:, :shift_idx]], dim=1) #enhance-a-video - if feta_args is not None: - if feta_start_percent <= current_step_percentage <= feta_end_percent: - enable_enhance() - else: - disable_enhance() + if feta_args is not None and feta_start_percent <= current_step_percentage <= feta_end_percent: + enable_enhance() + else: + disable_enhance() + #flow-edit if flowedit_args is not None: sigma = t / 1000.0 @@ -3103,6 +3228,9 @@ class WanVideoSampler: callback(idx, callback_latent, None, steps) else: pbar.update(1) + + if phantom_latents is not None: + x0 = x0[:,:-phantom_latents.shape[1]] if teacache_args is not None: states = transformer.teacache_state.states @@ -3362,6 +3490,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoVACEEncode": WanVideoVACEEncode, "WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame, "WanVideoVACEModelSelect": WanVideoVACEModelSelect, + "WanVideoPhantomEmbeds": WanVideoPhantomEmbeds, } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3397,4 +3526,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoVACEEncode": "WanVideo VACE Encode", "WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame", "WanVideoVACEModelSelect": "WanVideo VACE Model Select", + "WanVideoPhantomEmbeds": "WanVideo Phantom Embeds", } diff --git a/skyreels/nodes.py b/skyreels/nodes.py index 8a301f2..91f37ac 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -12,6 +12,8 @@ from ..wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler from ..nodes import optimized_scale from einops import rearrange +from ..enhance_a_video.globals import disable_enhance + import comfy.model_management as mm from comfy.utils import load_torch_file, ProgressBar, common_upscale from comfy.clip_vision import clip_preprocess, ClipVisionModel @@ -139,7 +141,7 @@ class WanVideoDiffusionForcingSampler: patcher = model model = model.model transformer = model.diffusion_model - + dtype = model["dtype"] device = mm.get_torch_device() offload_device = mm.unet_offload_device() @@ -305,6 +307,7 @@ class WanVideoDiffusionForcingSampler: "end_percent": unianimate_poses["end_percent"] } + disable_enhance() #not sure if this can work, disabling for now to avoid errors if it's enabled by another sampler freqs = None transformer.rope_embedder.k = None @@ -371,6 +374,7 @@ class WanVideoDiffusionForcingSampler: transformer.rel_l1_thresh = teacache_args["rel_l1_thresh"] transformer.teacache_start_step = teacache_args["start_step"] transformer.teacache_cache_device = teacache_args["cache_device"] + log.info(f"TeaCache: Using cache device: {transformer.teacache_state.cache_device}") transformer.teacache_end_step = len(init_timesteps)-1 if teacache_args["end_step"] == -1 else teacache_args["end_step"] transformer.teacache_use_coefficients = teacache_args["use_coefficients"] transformer.teacache_mode = teacache_args["mode"] @@ -410,7 +414,7 @@ class WanVideoDiffusionForcingSampler: #region model pred def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, vace_data=None, unianim_data=None, teacache_state=None): - with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True): + with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): if use_cfg_zero_star and (idx <= zero_star_steps) and use_zero_init: return latent_model_input*0, None @@ -525,7 +529,7 @@ class WanVideoDiffusionForcingSampler: #print("timestep", timestep) noise_pred, self.teacache_state = predict_with_cfg( - latent_model_input, + latent_model_input.to(dtype), cfg[i], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 96fcea5..923a9e5 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -196,9 +196,9 @@ def attention( elif attention_mode == 'sageattn': attn_mask = None - q = q.transpose(1, 2).to(dtype) - k = k.transpose(1, 2).to(dtype) - v = v.transpose(1, 2).to(dtype) + q = q.transpose(1, 2) + k = k.transpose(1, 2) + v = v.transpose(1, 2) out = sageattn_func( q, k, v, attn_mask=attn_mask, is_causal=causal, dropout_p=dropout_p) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 478ab25..dabc271 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -134,10 +134,10 @@ class WanRMSNorm(nn.Module): Args: x(Tensor): Shape [B, L, C] """ - return self._norm(x.float()).type_as(x) * self.weight + return self._norm(x)* self.weight def _norm(self, x): - return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) + return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps).to(x.dtype) class WanLayerNorm(nn.LayerNorm): @@ -150,7 +150,7 @@ class WanLayerNorm(nn.LayerNorm): Args: x(Tensor): Shape [B, L, C] """ - return super().forward(x.float()).type_as(x) + return super().forward(x) class WanSelfAttention(nn.Module): @@ -442,6 +442,20 @@ class WanAttentionBlock(nn.Module): # modulation self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5) + @torch.compiler.disable() + def get_mod(self, e): + if e.dim() == 3: + modulation = self.modulation # 1, 6, dim + e = (modulation.to(e.device) + e).chunk(6, dim=1) + elif e.dim() == 4: + modulation = self.modulation.unsqueeze(2) # 1, 6, 1, dim + e = (modulation.to(e.device) + e).chunk(6, dim=1) + e = [ei.squeeze(1) for ei in e] + return e + + def modulate(self, x, e): + return x * (1 + e[1]) + e[0] + def forward( self, x, @@ -467,16 +481,9 @@ class WanAttentionBlock(nn.Module): freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2] """ #e = (self.modulation.to(e.device) + e).chunk(6, dim=1) - - if e.dim() == 3: - modulation = self.modulation # 1, 6, dim - e = (modulation.to(e.device) + e).chunk(6, dim=1) - elif e.dim() == 4: - modulation = self.modulation.unsqueeze(2) # 1, 6, 1, dim - e = (modulation.to(e.device) + e).chunk(6, dim=1) - e = [ei.squeeze(1) for ei in e] + e = self.get_mod(e) - input_x = self.norm1(x) * (1 + e[1]) + e[0] + input_x = self.modulate(self.norm1(x), e) if camera_embed is not None: # encode ReCamMaster camera @@ -506,20 +513,23 @@ class WanAttentionBlock(nn.Module): if camera_embed is not None: y = self.projector(y) - x = x.to(torch.float32) + (y.to(torch.float32) * e[2].to(torch.float32)) + del input_x + + x = x + (y * e[2]) + del y # cross-attention & ffn function if (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and x.shape[0] == 1: x = self.split_cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) else: x = self.cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes) - + del e return x - + @torch.compiler.disable() def cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None): x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed) - y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3]) - x = x.to(torch.float32) + (y.to(torch.float32) * e[5]) + y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) + x = x + (y * e[5]) return x @torch.compiler.disable() @@ -574,9 +584,9 @@ class WanAttentionBlock(nn.Module): # Continue with FFN x = x + x_combined - y = self.ffn(self.norm2(x).float() * (1 + e[4]) + e[3]) - x = x.to(torch.float32) + (y.to(torch.float32) * e[5].to(torch.float32)) - return x + y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) + x = x + (y * e[5]) + return x class VaceWanAttentionBlock(WanAttentionBlock): def __init__( @@ -659,6 +669,16 @@ class Head(nn.Module): # modulation self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5) + def get_mod(self, e): + if e.dim() == 2: + modulation = self.modulation.to(e.device) # 1, 2, dim + e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) + elif e.dim() == 3: + modulation = self.modulation.to(e.device).unsqueeze(2) # 1, 2, seq, dim + e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) + e = [ei.squeeze(1) for ei in e] + return e + def forward(self, x, e): r""" Args: @@ -670,13 +690,7 @@ class Head(nn.Module): # normed = self.norm(x) # x = self.head(normed * (1 + e[1]) + e[0]) - if e.dim() == 2: - modulation = self.modulation.to(e.device) # 1, 2, dim - e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) - elif e.dim() == 3: - modulation = self.modulation.to(e.device).unsqueeze(2) # 1, 2, seq, dim - e = (modulation + e.unsqueeze(1)).chunk(2, dim=1) - e = [ei.squeeze(1) for ei in e] + e = self.get_mod(e) x = self.head(self.norm(x) * (1 + e[1]) + e[0]) return x @@ -734,6 +748,8 @@ class WanModel(ModelMixin, ConfigMixin): vace_layers=None, vace_in_dim=None, inject_sample_info=False, + add_ref_conv=False, + in_dim_ref_conv=16, ): r""" Initialize the diffusion model backbone. @@ -885,9 +901,15 @@ class WanModel(ModelMixin, ConfigMixin): if model_type == 'i2v' or model_type == 'fl2v': self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v') + #skyreels v2 if inject_sample_info: self.fps_embedding = nn.Embedding(2, dim) self.fps_projection = nn.Sequential(nn.Linear(dim, dim), nn.SiLU(), nn.Linear(dim, dim * 6)) + #fun 1.1 + if add_ref_conv: + self.ref_conv = nn.Conv2d(in_dim_ref_conv, dim, kernel_size=patch_size[1:], stride=patch_size[1:]) + else: + self.ref_conv = None def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False, vace_blocks_to_swap=None): log.info(f"Swapping {blocks_to_swap + 1} transformer blocks") @@ -944,7 +966,7 @@ class WanModel(ModelMixin, ConfigMixin): kwargs ): # embeddings - c = [self.vace_patch_embedding(u.unsqueeze(0)) for u in vace_context] + c = [self.vace_patch_embedding(u.unsqueeze(0).float()).to(x.dtype) for u in vace_context] c = [u.flatten(2).transpose(1, 2) for u in c] c = torch.cat([ torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], @@ -991,6 +1013,7 @@ class WanModel(ModelMixin, ConfigMixin): camera_embed=None, unianim_data=None, fps_embeds=None, + fun_ref = None ): r""" Forward pass through the diffusion model @@ -1032,13 +1055,13 @@ class WanModel(ModelMixin, ConfigMixin): if control_lora_enabled: self.expanded_patch_embedding.to(device) x = [ - self.expanded_patch_embedding(u.unsqueeze(0)) + self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x ] else: self.original_patch_embedding.to(self.main_device) x = [ - self.original_patch_embedding(u.unsqueeze(0)) + self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x ] @@ -1046,6 +1069,14 @@ class WanModel(ModelMixin, ConfigMixin): [torch.tensor(u.shape[2:], dtype=torch.long) for u in x]) x = [u.flatten(2).transpose(1, 2) for u in x] + + if self.ref_conv is not None and fun_ref is not None: + fun_ref = self.ref_conv(fun_ref).flatten(2).transpose(1, 2) + grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + seq_len += fun_ref.size(1) + F += 1 + x = [torch.concat([_fun_ref.unsqueeze(0), u], dim=1) for _fun_ref, u in zip(fun_ref, x)] + seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long) assert seq_lens.max() <= seq_len x = torch.cat([ @@ -1069,39 +1100,43 @@ class WanModel(ModelMixin, ConfigMixin): rope_func = "default" # time embeddings - with torch.autocast(device_type='cuda', dtype=torch.float32): - # e = self.time_embedding( - # sinusoidal_embedding_1d(self.freq_dim, t).float()) - # e0 = self.time_projection(e).unflatten(1, (6, self.dim)) - # assert e.dtype == torch.float32 and e0.dtype == torch.float32 - if t.dim() == 2: - b, f = t.shape - _flag_df = True - else: - _flag_df = False + + # e = self.time_embedding( + # sinusoidal_embedding_1d(self.freq_dim, t).float()) + # e0 = self.time_projection(e).unflatten(1, (6, self.dim)) + # assert e.dtype == torch.float32 and e0.dtype == torch.float32 + if t.dim() == 2: + b, f = t.shape + _flag_df = True + else: + _flag_df = False - e = self.time_embedding( - sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(self.patch_embedding.weight.dtype) - ) # b, dim - e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim + e = self.time_embedding( + sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(x.dtype) + ) # b, dim + e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim - if fps_embeds is not None: - fps_embeds = torch.tensor(fps_embeds, dtype=torch.long, device=device) - - fps_emb = self.fps_embedding(fps_embeds).float() - if _flag_df: - e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)).repeat(t.shape[1], 1, 1) - else: - e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)) + if fps_embeds is not None: + fps_embeds = torch.tensor(fps_embeds, dtype=torch.long, device=device) + fps_emb = self.fps_embedding(fps_embeds).to(e0.dtype) if _flag_df: - e = e.view(b, f, 1, 1, self.dim) - e0 = e0.view(b, f, 1, 1, 6, self.dim) - e = e.repeat(1, 1, grid_sizes[0][1], grid_sizes[0][2], 1).flatten(1, 3) - e0 = e0.repeat(1, 1, grid_sizes[0][1], grid_sizes[0][2], 1, 1).flatten(1, 3) - e0 = e0.transpose(1, 2).contiguous() + e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)).repeat(t.shape[1], 1, 1) + else: + e0 = e0 + self.fps_projection(fps_emb).unflatten(1, (6, self.dim)) - assert e.dtype == torch.float32 and e0.dtype == torch.float32 + if _flag_df: + e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim) + e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim) + + e = e.flatten(1, 3) + e0 = e0.flatten(1, 3) + + e0 = e0.transpose(1, 2) + if not e0.is_contiguous(): + e0 = e0.contiguous() + + e = e.to(self.offload_device, non_blocking=self.use_non_blocking) # context context_lens = None @@ -1112,7 +1147,7 @@ class WanModel(ModelMixin, ConfigMixin): torch.cat( [u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context - ])) + ]).to(x.dtype)) if self.offload_txt_emb: self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -1143,10 +1178,13 @@ class WanModel(ModelMixin, ConfigMixin): if self.teacache_use_coefficients: rescale_func = np.poly1d(self.teacache_coefficients[self.teacache_mode]) temb = e if self.teacache_mode == 'e' else e0 - accumulated_rel_l1_distance += rescale_func(((temb-previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean()).cpu().item()) + accumulated_rel_l1_distance += rescale_func(( + (temb.to(device) - previous_modulated_input).abs().mean() / previous_modulated_input.abs().mean() + ).cpu().item()) else: temb_relative_l1 = relative_l1_distance(previous_modulated_input, e0) accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(e0.device) + temb_relative_l1 + del temb #print("accumulated_rel_l1_distance", accumulated_rel_l1_distance) @@ -1155,8 +1193,10 @@ class WanModel(ModelMixin, ConfigMixin): else: should_calc = True accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) + accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) previous_modulated_input = e.clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.clone() + previous_modulated_input = previous_modulated_input.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) if not should_calc: x = x.to(previous_residual.dtype) + previous_residual.to(x.device) #log.info(f"TeaCache: Skipping uncond step {current_step+1}") @@ -1174,7 +1214,6 @@ class WanModel(ModelMixin, ConfigMixin): if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']: dwpose_emb = unianim_data['dwpose'] x += dwpose_emb * unianim_data['strength'] - # arguments kwargs = dict( e=e0, @@ -1198,11 +1237,11 @@ class WanModel(ModelMixin, ConfigMixin): if (data["start"] <= current_step_percentage <= data["end"]) or \ (data["end"] > 0 and current_step == 0 and current_step_percentage >= data["start"]): - vace_hints = self.forward_vace(x.to(torch.float32), data["context"], data["seq_len"], kwargs) + vace_hints = self.forward_vace(x, data["context"], data["seq_len"], kwargs) vace_hint_list.append(vace_hints) vace_scale_list.append(data["scale"]) else: - vace_hints = self.forward_vace(x.to(torch.float32), vace_data, seq_len, kwargs) + vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs) vace_hint_list.append(vace_hints) vace_scale_list.append(1.0) @@ -1216,7 +1255,7 @@ class WanModel(ModelMixin, ConfigMixin): continue if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.main_device) - x = block(x.to(torch.float32), **kwargs) + x = block(x, **kwargs) if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -1224,10 +1263,16 @@ class WanModel(ModelMixin, ConfigMixin): self.teacache_state.update( pred_id, previous_residual=(x.to(original_x.device) - original_x), - accumulated_rel_l1_distance=accumulated_rel_l1_distance.to(self.teacache_cache_device, non_blocking=self.use_non_blocking), - previous_modulated_input=previous_modulated_input.to(self.teacache_cache_device, non_blocking=self.use_non_blocking) + accumulated_rel_l1_distance=accumulated_rel_l1_distance, + previous_modulated_input=previous_modulated_input ) - x = self.head(x, e) + + if self.ref_conv is not None and fun_ref is not None: + full_ref_length = fun_ref.size(1) + x = x[:, full_ref_length:] + grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) + + x = self.head(x, e.to(x.device)) x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type] x = [u.float() for u in x] return (x, pred_id) if pred_id is not None else (x, None) @@ -1260,7 +1305,6 @@ class WanModel(ModelMixin, ConfigMixin): class TeaCacheState: def __init__(self, cache_device='cpu'): self.cache_device = cache_device - log.info(f"TeaCache: Using cache device: {self.cache_device}") self.states = {} self._next_pred_id = 0 @@ -1296,7 +1340,7 @@ class TeaCacheState: del self.states[pred_id] def clear_all(self): - self.states.clear() + self.states = {} self._next_pred_id = 0 def relative_l1_distance(last_tensor, current_tensor): @@ -1304,3 +1348,7 @@ def relative_l1_distance(last_tensor, current_tensor): norm = torch.abs(last_tensor).mean() relative_l1_distance = l1_distance / norm return relative_l1_distance.to(torch.float32).to(current_tensor.device) + +def get_tensor_memory(tensor): + memory_bytes = tensor.element_size() * tensor.nelement() + return f"{memory_bytes / (1024 * 1024):.2f} MB" \ No newline at end of file