From ad43eed0d030c84c233203706b01e0f80fc2f071 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 13 Jun 2025 12:31:13 +0300 Subject: [PATCH] Add MagCache and consolidate cache_args to support both Not really getting great results yet with MagCache, but at least pretty much on bar with TeaCache so it can be an option that hopefully improves in time --- .../wanvideo_1_3B_ReCamMaster_example_01.json | 8 +- .../wanvideo_1_3B_VACE_examples_03.json | 24 +- ...wanvideo_1_3B_control_lora_example_01.json | 8 +- ...wanvideo_480p_I2V_endframe_example_01.json | 8 +- .../wanvideo_480p_I2V_example_02.json | 8 +- .../wanvideo_ATI_testing_01.json | 4 +- .../wanvideo_FLF2V_720P_example_01.json | 8 +- ...anvideo_Fun_control_camera_example_01.json | 8 +- .../wanvideo_Fun_control_example_01.json | 8 +- ...anvideo_I2V_FantasyTalking_example_01.json | 8 +- .../wanvideo_T2V_example_02.json | 8 +- .../wanvideo_flowedit_I2V_example_01.json | 8 +- .../wanvideo_long_T2V_example_01.json | 8 +- ...nvideo_phantom_subject2vid_example_01.json | 8 +- .../wanvideo_skyreels_a2_example_01.json | 8 +- ...iffusion_forcing_extension_example_01.json | 634 +++++++++--------- .../wanvideo_vid2vid_example_01.json | 8 +- nodes.py | 209 +++--- skyreels/nodes.py | 42 +- wanvideo/modules/model.py | 139 +++- 20 files changed, 661 insertions(+), 503 deletions(-) diff --git a/example_workflows/wanvideo_1_3B_ReCamMaster_example_01.json b/example_workflows/wanvideo_1_3B_ReCamMaster_example_01.json index 78da4a8..3baa3f9 100644 --- a/example_workflows/wanvideo_1_3B_ReCamMaster_example_01.json +++ b/example_workflows/wanvideo_1_3B_ReCamMaster_example_01.json @@ -1587,9 +1587,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 335 }, { @@ -1668,8 +1668,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 335 ] diff --git a/example_workflows/wanvideo_1_3B_VACE_examples_03.json b/example_workflows/wanvideo_1_3B_VACE_examples_03.json index 9db57c8..79d925f 100644 --- a/example_workflows/wanvideo_1_3B_VACE_examples_03.json +++ b/example_workflows/wanvideo_1_3B_VACE_examples_03.json @@ -1751,8 +1751,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 334 ] @@ -1789,8 +1789,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 335 ] @@ -2252,8 +2252,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 350 ] @@ -2840,9 +2840,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 350 }, { @@ -2966,9 +2966,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 335 }, { @@ -5220,9 +5220,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 334 }, { diff --git a/example_workflows/wanvideo_1_3B_control_lora_example_01.json b/example_workflows/wanvideo_1_3B_control_lora_example_01.json index 9d474e9..6d8eeb6 100644 --- a/example_workflows/wanvideo_1_3B_control_lora_example_01.json +++ b/example_workflows/wanvideo_1_3B_control_lora_example_01.json @@ -176,8 +176,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 103 ] @@ -826,9 +826,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 103 }, { diff --git a/example_workflows/wanvideo_480p_I2V_endframe_example_01.json b/example_workflows/wanvideo_480p_I2V_endframe_example_01.json index d2fdaa0..3d42020 100644 --- a/example_workflows/wanvideo_480p_I2V_endframe_example_01.json +++ b/example_workflows/wanvideo_480p_I2V_endframe_example_01.json @@ -1666,8 +1666,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 106 ] @@ -1771,9 +1771,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 106 }, { diff --git a/example_workflows/wanvideo_480p_I2V_example_02.json b/example_workflows/wanvideo_480p_I2V_example_02.json index 56c008e..06b1bc2 100644 --- a/example_workflows/wanvideo_480p_I2V_example_02.json +++ b/example_workflows/wanvideo_480p_I2V_example_02.json @@ -365,8 +365,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 56 ] @@ -796,9 +796,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 56 }, { diff --git a/example_workflows/wanvideo_ATI_testing_01.json b/example_workflows/wanvideo_ATI_testing_01.json index d1027d6..768f3da 100644 --- a/example_workflows/wanvideo_ATI_testing_01.json +++ b/example_workflows/wanvideo_ATI_testing_01.json @@ -907,9 +907,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": null }, { diff --git a/example_workflows/wanvideo_FLF2V_720P_example_01.json b/example_workflows/wanvideo_FLF2V_720P_example_01.json index d19bd56..65b0cd9 100644 --- a/example_workflows/wanvideo_FLF2V_720P_example_01.json +++ b/example_workflows/wanvideo_FLF2V_720P_example_01.json @@ -1519,8 +1519,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 106 ] @@ -1589,9 +1589,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 106 }, { diff --git a/example_workflows/wanvideo_Fun_control_camera_example_01.json b/example_workflows/wanvideo_Fun_control_camera_example_01.json index 5e100e4..97c7296 100644 --- a/example_workflows/wanvideo_Fun_control_camera_example_01.json +++ b/example_workflows/wanvideo_Fun_control_camera_example_01.json @@ -747,8 +747,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 56 ] @@ -1263,9 +1263,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 56 }, { diff --git a/example_workflows/wanvideo_Fun_control_example_01.json b/example_workflows/wanvideo_Fun_control_example_01.json index fad03e7..ad8215a 100644 --- a/example_workflows/wanvideo_Fun_control_example_01.json +++ b/example_workflows/wanvideo_Fun_control_example_01.json @@ -794,8 +794,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 56 ] @@ -1719,9 +1719,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 56 }, { diff --git a/example_workflows/wanvideo_I2V_FantasyTalking_example_01.json b/example_workflows/wanvideo_I2V_FantasyTalking_example_01.json index 629156a..eedf2de 100644 --- a/example_workflows/wanvideo_I2V_FantasyTalking_example_01.json +++ b/example_workflows/wanvideo_I2V_FantasyTalking_example_01.json @@ -485,9 +485,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 89 }, { @@ -1699,8 +1699,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 89 ] diff --git a/example_workflows/wanvideo_T2V_example_02.json b/example_workflows/wanvideo_T2V_example_02.json index 124f772..79a2ebf 100644 --- a/example_workflows/wanvideo_T2V_example_02.json +++ b/example_workflows/wanvideo_T2V_example_02.json @@ -401,8 +401,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 56 ] @@ -878,9 +878,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 56 }, { diff --git a/example_workflows/wanvideo_flowedit_I2V_example_01.json b/example_workflows/wanvideo_flowedit_I2V_example_01.json index 6b48698..18f16af 100644 --- a/example_workflows/wanvideo_flowedit_I2V_example_01.json +++ b/example_workflows/wanvideo_flowedit_I2V_example_01.json @@ -1246,8 +1246,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 91 ] @@ -1313,9 +1313,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 91 }, { diff --git a/example_workflows/wanvideo_long_T2V_example_01.json b/example_workflows/wanvideo_long_T2V_example_01.json index 5188e4a..6f07fa6 100644 --- a/example_workflows/wanvideo_long_T2V_example_01.json +++ b/example_workflows/wanvideo_long_T2V_example_01.json @@ -402,8 +402,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 58 ] @@ -515,9 +515,9 @@ "link": 57 }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 58 }, { diff --git a/example_workflows/wanvideo_phantom_subject2vid_example_01.json b/example_workflows/wanvideo_phantom_subject2vid_example_01.json index 3a90c46..83fc469 100644 --- a/example_workflows/wanvideo_phantom_subject2vid_example_01.json +++ b/example_workflows/wanvideo_phantom_subject2vid_example_01.json @@ -451,8 +451,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 71 ] @@ -1121,9 +1121,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 71 }, { diff --git a/example_workflows/wanvideo_skyreels_a2_example_01.json b/example_workflows/wanvideo_skyreels_a2_example_01.json index c9feaab..49ab7bb 100644 --- a/example_workflows/wanvideo_skyreels_a2_example_01.json +++ b/example_workflows/wanvideo_skyreels_a2_example_01.json @@ -587,8 +587,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 56 ] @@ -3788,9 +3788,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 56 }, { diff --git a/example_workflows/wanvideo_skyreels_diffusion_forcing_extension_example_01.json b/example_workflows/wanvideo_skyreels_diffusion_forcing_extension_example_01.json index ec9b106..86c9e87 100644 --- a/example_workflows/wanvideo_skyreels_diffusion_forcing_extension_example_01.json +++ b/example_workflows/wanvideo_skyreels_diffusion_forcing_extension_example_01.json @@ -2,7 +2,7 @@ "id": "206247b6-9fec-4ed2-8927-e4f388c674d4", "revision": 0, "last_node_id": 196, - "last_link_id": 303, + "last_link_id": 306, "nodes": [ { "id": 42, @@ -1012,8 +1012,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "links": [ 195 ] @@ -1051,8 +1051,8 @@ "mode": 0, "inputs": [ { - "name": "TEACACHEARGS", - "type": "TEACACHEARGS", + "name": "CACHEARGS", + "type": "CACHEARGS", "link": 195 } ], @@ -1109,38 +1109,6 @@ "ExpArgs" ] }, - { - "id": 140, - "type": "GetNode", - "pos": [ - 724.8881225585938, - -579.0637817382812 - ], - "size": [ - 210, - 34 - ], - "flags": { - "collapsed": true - }, - "order": 18, - "mode": 0, - "inputs": [], - "outputs": [ - { - "name": "TEACACHEARGS", - "type": "TEACACHEARGS", - "links": [ - 235 - ] - } - ], - "title": "Get_TeaCache", - "properties": {}, - "widgets_values": [ - "TeaCache" - ] - }, { "id": 141, "type": "GetNode", @@ -1155,7 +1123,7 @@ "flags": { "collapsed": true }, - "order": 19, + "order": 18, "mode": 0, "inputs": [], "outputs": [ @@ -1185,7 +1153,7 @@ 226 ], "flags": {}, - "order": 20, + "order": 19, "mode": 0, "inputs": [], "outputs": [ @@ -1225,7 +1193,7 @@ 106 ], "flags": {}, - "order": 21, + "order": 20, "mode": 0, "inputs": [], "outputs": [ @@ -1300,7 +1268,7 @@ "flags": { "collapsed": true }, - "order": 22, + "order": 21, "mode": 0, "inputs": [], "outputs": [ @@ -1332,15 +1300,16 @@ "flags": { "collapsed": true }, - "order": 23, + "order": 22, "mode": 0, "inputs": [], "outputs": [ { - "name": "TEACACHEARGS", - "type": "TEACACHEARGS", + "name": "CACHEARGS", + "type": "CACHEARGS", "links": [ - 196 + 196, + 305 ] } ], @@ -1364,7 +1333,7 @@ "flags": { "collapsed": true }, - "order": 24, + "order": 23, "mode": 0, "inputs": [], "outputs": [ @@ -1396,7 +1365,7 @@ "flags": { "collapsed": true }, - "order": 25, + "order": 24, "mode": 0, "inputs": [], "outputs": [ @@ -1423,7 +1392,7 @@ ], "size": [ 428.4000244140625, - 860.4000244140625 + 880.4000244140625 ], "flags": {}, "order": 88, @@ -1457,10 +1426,10 @@ "link": 183 }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", - "link": 235 + "type": "CACHEARGS", + "link": 304 }, { "name": "slg_args", @@ -1506,7 +1475,8 @@ true, "unipc", 1, - "comfy" + "comfy", + "" ] }, { @@ -1521,7 +1491,7 @@ 154 ], "flags": {}, - "order": 26, + "order": 25, "mode": 0, "inputs": [], "outputs": [ @@ -1559,7 +1529,7 @@ 58 ], "flags": {}, - "order": 27, + "order": 26, "mode": 0, "inputs": [], "outputs": [ @@ -1594,7 +1564,7 @@ "flags": { "collapsed": true }, - "order": 28, + "order": 27, "mode": 0, "inputs": [], "outputs": [ @@ -1686,7 +1656,7 @@ "flags": { "collapsed": true }, - "order": 29, + "order": 28, "mode": 0, "inputs": [], "outputs": [ @@ -1771,7 +1741,7 @@ "flags": { "collapsed": true }, - "order": 30, + "order": 29, "mode": 0, "inputs": [], "outputs": [ @@ -1805,7 +1775,7 @@ "flags": { "collapsed": true }, - "order": 31, + "order": 30, "mode": 0, "inputs": [], "outputs": [ @@ -1839,7 +1809,7 @@ "flags": { "collapsed": true }, - "order": 32, + "order": 31, "mode": 0, "inputs": [], "outputs": [ @@ -1873,15 +1843,16 @@ "flags": { "collapsed": true }, - "order": 33, + "order": 32, "mode": 0, "inputs": [], "outputs": [ { - "name": "TEACACHEARGS", - "type": "TEACACHEARGS", + "name": "CACHEARGS", + "type": "CACHEARGS", "links": [ - 260 + 260, + 306 ] } ], @@ -1905,7 +1876,7 @@ "flags": { "collapsed": true }, - "order": 34, + "order": 33, "mode": 0, "inputs": [], "outputs": [ @@ -1937,7 +1908,7 @@ "flags": { "collapsed": true }, - "order": 35, + "order": 34, "mode": 0, "inputs": [], "outputs": [ @@ -1955,101 +1926,6 @@ "ExpArgs" ] }, - { - "id": 165, - "type": "WanVideoDiffusionForcingSampler", - "pos": [ - 5483.89599609375, - -510.88037109375 - ], - "size": [ - 428.4000244140625, - 860.4000244140625 - ], - "flags": {}, - "order": 103, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "WANVIDEOMODEL", - "link": 256 - }, - { - "name": "text_embeds", - "type": "WANVIDEOTEXTEMBEDS", - "link": 296 - }, - { - "name": "image_embeds", - "type": "WANVIDIMAGE_EMBEDS", - "link": 258 - }, - { - "name": "samples", - "shape": 7, - "type": "LATENT", - "link": null - }, - { - "name": "prefix_samples", - "shape": 7, - "type": "LATENT", - "link": 259 - }, - { - "name": "teacache_args", - "shape": 7, - "type": "TEACACHEARGS", - "link": 260 - }, - { - "name": "slg_args", - "shape": 7, - "type": "SLGARGS", - "link": 261 - }, - { - "name": "experimental_args", - "shape": 7, - "type": "EXPERIMENTALARGS", - "link": 262 - }, - { - "name": "unianimate_poses", - "shape": 7, - "type": "UNIANIMATE_POSE", - "link": null - } - ], - "outputs": [ - { - "name": "samples", - "type": "LATENT", - "links": [ - 252 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-WanVideoWrapper", - "ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503", - "Node name for S&R": "WanVideoDiffusionForcingSampler" - }, - "widgets_values": [ - 10, - 24.000000000000004, - 30, - 4.000000000000001, - 5.000000000000001, - 0, - "fixed", - true, - "unipc", - 1, - "comfy" - ] - }, { "id": 153, "type": "WanVideoEncode", @@ -2130,6 +2006,7 @@ }, { "name": "negative", + "shape": 7, "type": "CONDITIONING", "link": 55 } @@ -2207,7 +2084,7 @@ 60 ], "flags": {}, - "order": 36, + "order": 35, "mode": 0, "inputs": [], "outputs": [ @@ -2343,7 +2220,7 @@ "flags": { "collapsed": true }, - "order": 37, + "order": 36, "mode": 0, "inputs": [], "outputs": [ @@ -2654,101 +2531,6 @@ 80 ] }, - { - "id": 104, - "type": "WanVideoDiffusionForcingSampler", - "pos": [ - 3010, - -400 - ], - "size": [ - 428.4000244140625, - 860.4000244140625 - ], - "flags": {}, - "order": 97, - "mode": 0, - "inputs": [ - { - "name": "model", - "type": "WANVIDEOMODEL", - "link": 190 - }, - { - "name": "text_embeds", - "type": "WANVIDEOTEXTEMBEDS", - "link": 192 - }, - { - "name": "image_embeds", - "type": "WANVIDIMAGE_EMBEDS", - "link": 207 - }, - { - "name": "samples", - "shape": 7, - "type": "LATENT", - "link": null - }, - { - "name": "prefix_samples", - "shape": 7, - "type": "LATENT", - "link": 180 - }, - { - "name": "teacache_args", - "shape": 7, - "type": "TEACACHEARGS", - "link": 196 - }, - { - "name": "slg_args", - "shape": 7, - "type": "SLGARGS", - "link": 240 - }, - { - "name": "experimental_args", - "shape": 7, - "type": "EXPERIMENTALARGS", - "link": 198 - }, - { - "name": "unianimate_poses", - "shape": 7, - "type": "UNIANIMATE_POSE", - "link": null - } - ], - "outputs": [ - { - "name": "samples", - "type": "LATENT", - "links": [ - 178 - ] - } - ], - "properties": { - "cnr_id": "ComfyUI-WanVideoWrapper", - "ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503", - "Node name for S&R": "WanVideoDiffusionForcingSampler" - }, - "widgets_values": [ - 24, - 24.000000000000004, - 30, - 4.000000000000001, - 5.000000000000001, - 0, - "fixed", - true, - "unipc", - 1, - "comfy" - ] - }, { "id": 90, "type": "VHS_VideoCombine", @@ -3137,7 +2919,7 @@ "flags": { "collapsed": true }, - "order": 38, + "order": 37, "mode": 0, "inputs": [], "outputs": [ @@ -3434,7 +3216,7 @@ 203.9819793701172 ], "flags": {}, - "order": 39, + "order": 38, "mode": 0, "inputs": [], "outputs": [ @@ -3474,7 +3256,7 @@ 139.27662658691406 ], "flags": {}, - "order": 40, + "order": 39, "mode": 0, "inputs": [], "outputs": [], @@ -3621,7 +3403,7 @@ "flags": { "collapsed": true }, - "order": 41, + "order": 40, "mode": 0, "inputs": [], "outputs": [ @@ -3663,6 +3445,7 @@ }, { "name": "negative", + "shape": 7, "type": "CONDITIONING", "link": 265 } @@ -3695,7 +3478,7 @@ 95.97175598144531 ], "flags": {}, - "order": 42, + "order": 41, "mode": 0, "inputs": [], "outputs": [], @@ -3718,7 +3501,7 @@ 106 ], "flags": {}, - "order": 43, + "order": 42, "mode": 0, "inputs": [], "outputs": [ @@ -3753,10 +3536,10 @@ ], "size": [ 528.6734619140625, - 234 + 254 ], "flags": {}, - "order": 44, + "order": 43, "mode": 0, "inputs": [ { @@ -3788,6 +3571,12 @@ "shape": 7, "type": "VACEPATH", "link": null + }, + { + "name": "fantasytalking_model", + "shape": 7, + "type": "FANTASYTALKINGMODEL", + "link": null } ], "outputs": [ @@ -3829,7 +3618,7 @@ "flags": { "collapsed": true }, - "order": 45, + "order": 44, "mode": 0, "inputs": [], "outputs": [ @@ -3977,7 +3766,7 @@ 88 ], "flags": {}, - "order": 46, + "order": 45, "mode": 0, "inputs": [], "outputs": [], @@ -4000,7 +3789,7 @@ 451.9747314453125 ], "flags": {}, - "order": 47, + "order": 46, "mode": 0, "inputs": [ { @@ -4090,6 +3879,12 @@ "type": "IMAGE", "link": 298 }, + { + "name": "get_image_size", + "shape": 7, + "type": "IMAGE", + "link": null + }, { "name": "width_input", "shape": 7, @@ -4101,12 +3896,6 @@ "shape": 7, "type": "INT", "link": null - }, - { - "name": "get_image_size", - "shape": 7, - "type": "IMAGE", - "link": null } ], "outputs": [ @@ -4158,7 +3947,7 @@ 85.25131225585938 ], "flags": {}, - "order": 48, + "order": 47, "mode": 0, "inputs": [], "outputs": [ @@ -4233,7 +4022,7 @@ 88 ], "flags": {}, - "order": 49, + "order": 48, "mode": 0, "inputs": [], "outputs": [], @@ -4371,7 +4160,7 @@ 314 ], "flags": {}, - "order": 50, + "order": 49, "mode": 0, "inputs": [], "outputs": [ @@ -4452,7 +4241,7 @@ 209.54696655273438 ], "flags": {}, - "order": 51, + "order": 50, "mode": 0, "inputs": [], "outputs": [], @@ -4524,7 +4313,7 @@ 130 ], "flags": {}, - "order": 52, + "order": 51, "mode": 4, "inputs": [], "outputs": [ @@ -4564,7 +4353,7 @@ "flags": { "collapsed": true }, - "order": 53, + "order": 52, "mode": 0, "inputs": [], "outputs": [ @@ -4658,7 +4447,7 @@ "flags": { "collapsed": true }, - "order": 54, + "order": 53, "mode": 0, "inputs": [], "outputs": [ @@ -4692,7 +4481,7 @@ "flags": { "collapsed": true }, - "order": 55, + "order": 54, "mode": 0, "inputs": [], "outputs": [ @@ -4724,7 +4513,7 @@ 166.29330444335938 ], "flags": {}, - "order": 56, + "order": 55, "mode": 0, "inputs": [], "outputs": [], @@ -4734,6 +4523,229 @@ ], "color": "#432", "bgcolor": "#653" + }, + { + "id": 140, + "type": "GetNode", + "pos": [ + 724.8883666992188, + -591.3699340820312 + ], + "size": [ + 210, + 50 + ], + "flags": { + "collapsed": true + }, + "order": 56, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "CACHEARGS", + "type": "CACHEARGS", + "links": [ + 235, + 304 + ] + } + ], + "title": "Get_TeaCache", + "properties": {}, + "widgets_values": [ + "TeaCache" + ] + }, + { + "id": 104, + "type": "WanVideoDiffusionForcingSampler", + "pos": [ + 3010, + -400 + ], + "size": [ + 428.4000244140625, + 860.4000244140625 + ], + "flags": {}, + "order": 97, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 190 + }, + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "link": 192 + }, + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 207 + }, + { + "name": "samples", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "prefix_samples", + "shape": 7, + "type": "LATENT", + "link": 180 + }, + { + "name": "cache_args", + "shape": 7, + "type": "CACHEARGS", + "link": 305 + }, + { + "name": "slg_args", + "shape": 7, + "type": "SLGARGS", + "link": 240 + }, + { + "name": "experimental_args", + "shape": 7, + "type": "EXPERIMENTALARGS", + "link": 198 + }, + { + "name": "unianimate_poses", + "shape": 7, + "type": "UNIANIMATE_POSE", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 178 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503", + "Node name for S&R": "WanVideoDiffusionForcingSampler" + }, + "widgets_values": [ + 24, + 24.000000000000004, + 30, + 4.000000000000001, + 5.000000000000001, + 0, + "fixed", + true, + "unipc", + 1, + "comfy" + ] + }, + { + "id": 165, + "type": "WanVideoDiffusionForcingSampler", + "pos": [ + 5483.89599609375, + -510.88037109375 + ], + "size": [ + 428.4000244140625, + 860.4000244140625 + ], + "flags": {}, + "order": 103, + "mode": 0, + "inputs": [ + { + "name": "model", + "type": "WANVIDEOMODEL", + "link": 256 + }, + { + "name": "text_embeds", + "type": "WANVIDEOTEXTEMBEDS", + "link": 296 + }, + { + "name": "image_embeds", + "type": "WANVIDIMAGE_EMBEDS", + "link": 258 + }, + { + "name": "samples", + "shape": 7, + "type": "LATENT", + "link": null + }, + { + "name": "prefix_samples", + "shape": 7, + "type": "LATENT", + "link": 259 + }, + { + "name": "cache_args", + "shape": 7, + "type": "CACHEARGS", + "link": 306 + }, + { + "name": "slg_args", + "shape": 7, + "type": "SLGARGS", + "link": 261 + }, + { + "name": "experimental_args", + "shape": 7, + "type": "EXPERIMENTALARGS", + "link": 262 + }, + { + "name": "unianimate_poses", + "shape": 7, + "type": "UNIANIMATE_POSE", + "link": null + } + ], + "outputs": [ + { + "name": "samples", + "type": "LATENT", + "links": [ + 252 + ] + } + ], + "properties": { + "cnr_id": "ComfyUI-WanVideoWrapper", + "ver": "e5a326c9811514f2c08c89bccea9a7c731d9a503", + "Node name for S&R": "WanVideoDiffusionForcingSampler" + }, + "widgets_values": [ + 10, + 24.000000000000004, + 30, + 4.000000000000001, + 5.000000000000001, + 0, + "fixed", + true, + "unipc", + 1, + "comfy" + ] } ], "links": [ @@ -4889,14 +4901,6 @@ 0, "*" ], - [ - 196, - 115, - 0, - 104, - 5, - "TEACACHEARGS" - ], [ 197, 87, @@ -5089,14 +5093,6 @@ 0, "IMAGE" ], - [ - 235, - 140, - 0, - 103, - 5, - "TEACACHEARGS" - ], [ 236, 141, @@ -5217,14 +5213,6 @@ 4, "LATENT" ], - [ - 260, - 162, - 0, - 165, - 5, - "TEACACHEARGS" - ], [ 261, 163, @@ -5456,6 +5444,30 @@ 156, 2, "INT" + ], + [ + 304, + 140, + 0, + 103, + 5, + "CACHEARGS" + ], + [ + 305, + 115, + 0, + 104, + 5, + "CACHEARGS" + ], + [ + 306, + 162, + 0, + 165, + 5, + "CACHEARGS" ] ], "groups": [ @@ -5528,13 +5540,13 @@ "config": {}, "extra": { "ds": { - "scale": 1.191817653772724, + "scale": 0.611590904484147, "offset": [ - 1695.7620823297345, - 1138.5391291690546 + 426.87167769967925, + 1142.3743330459465 ] }, - "frontendVersion": "1.17.3", + "frontendVersion": "1.22.0", "node_versions": { "ComfyUI-WanVideoWrapper": "5a2383621a05825d0d0437781afcb8552d9590fd", "comfy-core": "0.3.26", diff --git a/example_workflows/wanvideo_vid2vid_example_01.json b/example_workflows/wanvideo_vid2vid_example_01.json index fe6692d..e9e5973 100644 --- a/example_workflows/wanvideo_vid2vid_example_01.json +++ b/example_workflows/wanvideo_vid2vid_example_01.json @@ -668,8 +668,8 @@ "inputs": [], "outputs": [ { - "name": "teacache_args", - "type": "TEACACHEARGS", + "name": "cache_args", + "type": "CACHEARGS", "slot_index": 0, "links": [ 62 @@ -736,9 +736,9 @@ "link": null }, { - "name": "teacache_args", + "name": "cache_args", "shape": 7, - "type": "TEACACHEARGS", + "type": "CACHEARGS", "link": 62 }, { diff --git a/nodes.py b/nodes.py index e01c490..c663ad9 100644 --- a/nodes.py +++ b/nodes.py @@ -111,8 +111,8 @@ class WanVideoTeaCache: "mode": (["e", "e0"], {"default": "e", "tooltip": "Choice between using e (time embeds, default) or e0 (modulated time embeds)"}), }, } - RETURN_TYPES = ("TEACACHEARGS",) - RETURN_NAMES = ("teacache_args",) + RETURN_TYPES = ("CACHEARGS",) + RETURN_NAMES = ("cache_args",) FUNCTION = "process" CATEGORY = "WanVideoWrapper" DESCRIPTION = """ @@ -141,18 +141,54 @@ Official recommended values https://github.com/ali-vilab/TeaCache/tree/main/TeaC def process(self, rel_l1_thresh, start_step, end_step, cache_device, use_coefficients, mode="e"): if cache_device == "main_device": - teacache_device = mm.get_torch_device() + cache_device = mm.get_torch_device() else: - teacache_device = mm.unet_offload_device() - teacache_args = { + cache_device = mm.unet_offload_device() + cache_args = { + "cache_type": "TeaCache", "rel_l1_thresh": rel_l1_thresh, "start_step": start_step, "end_step": end_step, - "cache_device": teacache_device, + "cache_device": cache_device, "use_coefficients": use_coefficients, "mode": mode, } - return (teacache_args,) + return (cache_args,) + +class WanVideoMagCache: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "magcache_thresh": ("FLOAT", {"default": 0.02, "min": 0.0, "max": 0.3, "step": 0.001, "tooltip": "How strongly to cache the output of diffusion model. This value must be non-negative."}), + "magcache_K": ("INT", {"default": 4, "min": 0, "max": 6, "step": 1, "tooltip": "The maxium skip steps of MagCache."}), + "start_step": ("INT", {"default": 1, "min": 0, "max": 9999, "step": 1, "tooltip": "Step to start applying MagCache"}), + "end_step": ("INT", {"default": -1, "min": -1, "max": 9999, "step": 1, "tooltip": "Step to end applying MagCache"}), + "cache_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Device to cache to"}), + }, + } + RETURN_TYPES = ("CACHEARGS",) + RETURN_NAMES = ("cache_args",) + FUNCTION = "setargs" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + DESCRIPTION = "MagCache for WanVideoWrapper, source https://github.com/Zehong-Ma/MagCache" + + def setargs(self, magcache_thresh, magcache_K, start_step, end_step, cache_device): + if cache_device == "main_device": + cache_device = mm.get_torch_device() + else: + cache_device = mm.unet_offload_device() + + cache_args = { + "cache_type": "MagCache", + "magcache_thresh": magcache_thresh, + "magcache_K": magcache_K, + "start_step": start_step, + "end_step": end_step, + "cache_device": cache_device, + } + return (cache_args,) class WanVideoModel(comfy.model_base.BaseModel): @@ -625,6 +661,14 @@ class WanVideoModelLoader: "e0": [8.10705460e+03, 2.13393892e+03, -3.72934672e+02, 1.66203073e+01, -4.17769401e-02], }, } + + magcache_ratios_map = { + "1_3B": np.array([1.0]*2+[1.0124, 1.02213, 1.00166, 1.0041, 0.99791, 1.00061, 0.99682, 0.99762, 0.99634, 0.99685, 0.99567, 0.99586, 0.99416, 0.99422, 0.99578, 0.99575, 0.9957, 0.99563, 0.99511, 0.99506, 0.99535, 0.99531, 0.99552, 0.99549, 0.99541, 0.99539, 0.9954, 0.99536, 0.99489, 0.99485, 0.99518, 0.99514, 0.99484, 0.99478, 0.99481, 0.99479, 0.99415, 0.99413, 0.99419, 0.99416, 0.99396, 0.99393, 0.99388, 0.99386, 0.99349, 0.99349, 0.99309, 0.99304, 0.9927, 0.9927, 0.99228, 0.99226, 0.99171, 0.9917, 0.99137, 0.99135, 0.99068, 0.99063, 0.99005, 0.99003, 0.98944, 0.98942, 0.98849, 0.98849, 0.98758, 0.98757, 0.98644, 0.98643, 0.98504, 0.98503, 0.9836, 0.98359, 0.98202, 0.98201, 0.97977, 0.97978, 0.97717, 0.97718, 0.9741, 0.97411, 0.97003, 0.97002, 0.96538, 0.96541, 0.9593, 0.95933, 0.95086, 0.95089, 0.94013, 0.94019, 0.92402, 0.92414, 0.90241, 0.9026, 0.86821, 0.86868, 0.81838, 0.81939]), + "14B": np.array([1.0]*2+[1.02504, 1.03017, 1.00025, 1.00251, 0.9985, 0.99962, 0.99779, 0.99771, 0.9966, 0.99658, 0.99482, 0.99476, 0.99467, 0.99451, 0.99664, 0.99656, 0.99434, 0.99431, 0.99533, 0.99545, 0.99468, 0.99465, 0.99438, 0.99434, 0.99516, 0.99517, 0.99384, 0.9938, 0.99404, 0.99401, 0.99517, 0.99516, 0.99409, 0.99408, 0.99428, 0.99426, 0.99347, 0.99343, 0.99418, 0.99416, 0.99271, 0.99269, 0.99313, 0.99311, 0.99215, 0.99215, 0.99218, 0.99215, 0.99216, 0.99217, 0.99163, 0.99161, 0.99138, 0.99135, 0.98982, 0.9898, 0.98996, 0.98995, 0.9887, 0.98866, 0.98772, 0.9877, 0.98767, 0.98765, 0.98573, 0.9857, 0.98501, 0.98498, 0.9838, 0.98376, 0.98177, 0.98173, 0.98037, 0.98035, 0.97678, 0.97677, 0.97546, 0.97543, 0.97184, 0.97183, 0.96711, 0.96708, 0.96349, 0.96345, 0.95629, 0.95625, 0.94926, 0.94929, 0.93964, 0.93961, 0.92511, 0.92504, 0.90693, 0.90678, 0.8796, 0.87945, 0.86111, 0.86189]), + "i2v_480": np.array([1.0]*2+[0.98783, 0.98993, 0.97559, 0.97593, 0.98311, 0.98319, 0.98202, 0.98225, 0.9888, 0.98878, 0.98762, 0.98759, 0.98957, 0.98971, 0.99052, 0.99043, 0.99383, 0.99384, 0.98857, 0.9886, 0.99065, 0.99068, 0.98845, 0.98847, 0.99057, 0.99057, 0.98957, 0.98961, 0.98601, 0.9861, 0.98823, 0.98823, 0.98756, 0.98759, 0.98808, 0.98814, 0.98721, 0.98724, 0.98571, 0.98572, 0.98543, 0.98544, 0.98157, 0.98165, 0.98411, 0.98413, 0.97952, 0.97953, 0.98149, 0.9815, 0.9774, 0.97742, 0.97825, 0.97826, 0.97355, 0.97361, 0.97085, 0.97087, 0.97056, 0.97055, 0.96588, 0.96587, 0.96113, 0.96124, 0.9567, 0.95681, 0.94961, 0.94969, 0.93973, 0.93988, 0.93217, 0.93224, 0.91878, 0.91896, 0.90955, 0.90954, 0.92617, 0.92616]), + "i2v_720": np.array([1.0]*2+[0.99428, 0.99498, 0.98588, 0.98621, 0.98273, 0.98281, 0.99018, 0.99023, 0.98911, 0.98917, 0.98646, 0.98652, 0.99454, 0.99456, 0.9891, 0.98909, 0.99124, 0.99127, 0.99102, 0.99103, 0.99215, 0.99212, 0.99515, 0.99515, 0.99576, 0.99572, 0.99068, 0.99072, 0.99097, 0.99097, 0.99166, 0.99169, 0.99041, 0.99042, 0.99201, 0.99198, 0.99101, 0.99101, 0.98599, 0.98603, 0.98845, 0.98844, 0.98848, 0.98851, 0.98862, 0.98857, 0.98718, 0.98719, 0.98497, 0.98497, 0.98264, 0.98263, 0.98389, 0.98393, 0.97938, 0.9794, 0.97535, 0.97536, 0.97498, 0.97499, 0.973, 0.97301, 0.96827, 0.96828, 0.96261, 0.96263, 0.95335, 0.9534, 0.94649, 0.94655, 0.93397, 0.93414, 0.91636, 0.9165, 0.89088, 0.89109, 0.8679, 0.86768]), + } + model_variant = "14B" #default to this if model_type == "i2v" or model_type == "fl2v": if "480" in model or "fun" in model.lower() or "a2" in model.lower() or "540" in model: #just a guess for the Fun model for now... @@ -653,6 +697,7 @@ class WanVideoModelLoader: "main_device": device, "offload_device": offload_device, "teacache_coefficients": teacache_coefficients_map[model_variant], + "magcache_ratios": magcache_ratios_map[model_variant], "vace_layers": vace_layers, "vace_in_dim": vace_in_dim, "inject_sample_info": True if "fps_embedding.weight" in sd else False, @@ -2439,7 +2484,7 @@ class WanVideoSampler: "denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), "feta_args": ("FETAARGS", ), "context_options": ("WANVIDCONTEXT", ), - "teacache_args": ("TEACACHEARGS", ), + "cache_args": ("CACHEARGS", ), "flowedit_args": ("FLOWEDITARGS", ), "batched_cfg": ("BOOLEAN", {"default": False, "tooltip": "Batch cond and uncond for faster sampling, possibly faster on some hardware, uses more memory"}), "slg_args": ("SLGARGS", ), @@ -2460,9 +2505,9 @@ class WanVideoSampler: def process(self, model, text_embeds, image_embeds, shift, steps, cfg, seed, scheduler, riflex_freq_index, force_offload=True, samples=None, feta_args=None, denoise_strength=1.0, context_options=None, - teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, + cache_args=None, teacache_args=None, flowedit_args=None, batched_cfg=False, slg_args=None, rope_function="default", loop_args=None, experimental_args=None, sigmas=None, unianimate_poses=None, fantasytalking_embeds=None, uni3c_embeds=None): - #assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache." + patcher = model model = model.model transformer = model.diffusion_model @@ -2938,19 +2983,29 @@ class WanVideoSampler: feta_args = None disable_enhance() - # Initialize TeaCache if enabled - if teacache_args is not None: - transformer.enable_teacache = True - 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"] - transformer.teacache_state.clear_all() - else: - transformer.enable_teacache = False + # Initialize Cache if enabled + transformer.enable_teacache = transformer.enable_magcache = False + if teacache_args is not None: #for backward compatibility on old workflows + cache_args = teacache_args + if cache_args is not None: + transformer.cache_device = cache_args["cache_device"] + if cache_args["cache_type"] == "TeaCache": + log.info(f"TeaCache: Using cache device: {transformer.cache_device}") + transformer.teacache_state.clear_all() + transformer.enable_teacache = True + transformer.rel_l1_thresh = cache_args["rel_l1_thresh"] + transformer.teacache_start_step = cache_args["start_step"] + transformer.teacache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.teacache_use_coefficients = cache_args["use_coefficients"] + transformer.teacache_mode = cache_args["mode"] + elif cache_args["cache_type"] == "MagCache": + log.info(f"MagCache: Using cache device: {transformer.cache_device}") + transformer.magcache_state.clear_all() + transformer.enable_magcache = True + transformer.magcache_start_step = cache_args["start_step"] + transformer.magcache_end_step = len(timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.magcache_thresh = cache_args["magcache_thresh"] + transformer.magcache_K = cache_args["magcache_K"] if slg_args is not None: assert batched_cfg is not None, "Batched cfg is not supported with SLG" @@ -2960,12 +3015,12 @@ class WanVideoSampler: else: transformer.slg_blocks = None - self.teacache_state = [None, None] + self.cache_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 = [] + self.cache_state = [None, None, None] + self.cache_state_source = [None, None] + self.cache_states_context = [] if flowedit_args is not None: source_embeds = flowedit_args["source_embeds"] @@ -3022,7 +3077,7 @@ 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, audio_proj=None, control_camera_latents=None, add_cond=None, teacache_state=None): + control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, add_cond=None, cache_state=None): z = z.to(dtype) with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])): @@ -3087,8 +3142,8 @@ class WanVideoSampler: 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 teacache_state is not None and len(teacache_state) != 3: - teacache_state.append(None) + if cache_state is not None and len(cache_state) != 3: + cache_state.append(None) if not use_phantom: z_pos = z_neg = z @@ -3146,10 +3201,10 @@ class WanVideoSampler: if not batched_cfg: #cond - noise_pred_cond, teacache_state_cond = transformer( + noise_pred_cond, cache_state_cond = transformer( [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, + pred_id=cache_state[0] if cache_state else None, vace_data=vace_data, attn_cond=attn_cond, **base_params ) @@ -3162,44 +3217,44 @@ class WanVideoSampler: scale_high=fresca_scale_high, freq_cutoff=fresca_freq_cutoff, ) - return noise_pred_cond, [teacache_state_cond] + return noise_pred_cond, [cache_state_cond] #uncond if fantasytalking_embeds is not None: if not math.isclose(audio_cfg_scale[idx], 1.0): base_params['audio_proj'] = None - noise_pred_uncond, teacache_state_uncond = transformer( + noise_pred_uncond, cache_state_uncond = transformer( [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, + pred_id=cache_state[1] if cache_state else None, vace_data=vace_data, attn_cond=attn_cond_neg, **base_params ) noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device) #phantom if use_phantom and not math.isclose(phantom_cfg_scale[idx], 1.0): - noise_pred_phantom, teacache_state_phantom = transformer( + noise_pred_phantom, cache_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, + pred_id=cache_state[2] if cache_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[idx] * (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] + return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_phantom] #fantasytalking if fantasytalking_embeds is not None: if not math.isclose(audio_cfg_scale[idx], 1.0): - if teacache_state is not None and len(teacache_state) != 3: - teacache_state.append(None) + if cache_state is not None and len(cache_state) != 3: + cache_state.append(None) base_params['audio_proj'] = None - noise_pred_no_audio, teacache_state_audio = transformer( + noise_pred_no_audio, cache_state_audio = transformer( [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[2] if teacache_state else None, + pred_id=cache_state[2] if cache_state else None, vace_data=vace_data, **base_params ) @@ -3209,16 +3264,16 @@ class WanVideoSampler: + cfg_scale * (noise_pred_no_audio - noise_pred_uncond) + audio_cfg_scale[idx] * (noise_pred_cond - noise_pred_no_audio) ) - return noise_pred, [teacache_state_cond, teacache_state_uncond, teacache_state_audio] + return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_audio] #batched else: - teacache_state_uncond = None - [noise_pred_cond, noise_pred_uncond], teacache_state_cond = transformer( + cache_state_uncond = None + [noise_pred_cond, noise_pred_uncond], cache_state_cond = transformer( [z] + [z], context=positive_embeds + negative_embeds, y=[image_cond_input] + [image_cond_input] if image_cond_input is not None else None, clip_fea=clip_fea.repeat(2,1,1), is_uncond=False, current_step_percentage=current_step_percentage, - pred_id=teacache_state[0] if teacache_state else None, + pred_id=cache_state[0] if cache_state else None, **base_params ) #cfg @@ -3245,7 +3300,7 @@ class WanVideoSampler: noise_pred = noise_pred_uncond * alpha + cfg_scale * (noise_pred_cond - noise_pred_uncond * alpha) - return noise_pred, [teacache_state_cond, teacache_state_uncond] + return noise_pred, [cache_state_cond, cache_state_uncond] log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*8}x{latent.shape[2]*8} with {steps} steps") @@ -3332,8 +3387,8 @@ class WanVideoSampler: for c in context_queue: window_id = self.window_tracker.get_window_id(c) - if teacache_args is not None: - current_teacache = self.window_tracker.get_teacache(window_id, self.teacache_state) + if cache_args is not None: + current_teacache = self.window_tracker.get_teacache(window_id, self.cache_state) else: current_teacache = None @@ -3358,21 +3413,21 @@ class WanVideoSampler: timestep, idx, partial_img_emb, control_latents, source_clip_fea, current_teacache) - if teacache_args is not None: - self.window_tracker.teacache_states[window_id] = new_teacache + if cache_args is not None: + self.window_tracker.cache_states[window_id] = new_teacache window_mask = create_window_mask(vt_src_context, c, latent_video_length, context_overlap) vt_src[:, c, :, :] += vt_src_context * window_mask counter[:, c, :, :] += window_mask vt_src /= counter else: - vt_src, self.teacache_state_source = predict_with_cfg( + vt_src, self.cache_state_source = predict_with_cfg( zt_src, cfg[idx], source_embeds["prompt_embeds"], source_embeds["negative_prompt_embeds"], timestep, idx, source_image_cond, source_clip_fea, control_latents, - teacache_state=self.teacache_state_source) + cache_state=self.cache_state_source) else: if idx == len(timesteps) - drift_steps: x_tgt = zt_tgt @@ -3386,8 +3441,8 @@ class WanVideoSampler: for c in context_queue: window_id = self.window_tracker.get_window_id(c) - if teacache_args is not None: - current_teacache = self.window_tracker.get_teacache(window_id, self.teacache_state) + if cache_args is not None: + current_teacache = self.window_tracker.get_teacache(window_id, self.cache_state) else: current_teacache = None @@ -3415,20 +3470,20 @@ class WanVideoSampler: timestep, idx, partial_img_emb, partial_control_latents, clip_fea, current_teacache) - if teacache_args is not None: - self.window_tracker.teacache_states[window_id] = new_teacache + if cache_args is not None: + self.window_tracker.cache_states[window_id] = new_teacache window_mask = create_window_mask(vt_tgt_context, c, latent_video_length, context_overlap) vt_tgt[:, c, :, :] += vt_tgt_context * window_mask counter[:, c, :, :] += window_mask vt_tgt /= counter else: - vt_tgt, self.teacache_state = predict_with_cfg( + vt_tgt, self.cache_state = predict_with_cfg( zt_tgt, cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], timestep, idx, image_cond, clip_fea, control_latents, - teacache_state=self.teacache_state) + cache_state=self.cache_state) v_delta = vt_tgt - vt_src x_tgt = x_tgt.to(torch.float32) v_delta = v_delta.to(torch.float32) @@ -3443,8 +3498,8 @@ class WanVideoSampler: for c in context_queue: window_id = self.window_tracker.get_window_id(c) - if teacache_args is not None: - current_teacache = self.window_tracker.get_teacache(window_id, self.teacache_state) + if cache_args is not None: + current_teacache = self.window_tracker.get_teacache(window_id, self.cache_state) else: current_teacache = None @@ -3519,8 +3574,8 @@ class WanVideoSampler: partial_control_camera_latents, partial_add_cond, current_teacache) - if teacache_args is not None: - self.window_tracker.teacache_states[window_id] = new_teacache + if cache_args is not None: + self.window_tracker.cache_states[window_id] = new_teacache window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped) noise_pred[:, c] += noise_pred_context * window_mask @@ -3528,13 +3583,13 @@ class WanVideoSampler: noise_pred /= counter #region normal inference else: - noise_pred, self.teacache_state = predict_with_cfg( + noise_pred, self.cache_state = predict_with_cfg( latent_model_input, cfg[idx], text_embeds["prompt_embeds"], text_embeds["negative_prompt_embeds"], timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond, - teacache_state=self.teacache_state) + cache_state=self.cache_state) if latent_shift_loop: #reverse latent shift @@ -3580,8 +3635,9 @@ class WanVideoSampler: if phantom_latents is not None: x0 = x0[:,:-phantom_latents.shape[1]] - if teacache_args is not None: - states = transformer.teacache_state.states + if cache_args is not None: + cache_type = cache_args["cache_type"] + states = transformer.teacache_state.states if cache_type == "TeaCache" else transformer.magcache_state.states state_names = { 0: "conditional", 1: "unconditional" @@ -3589,13 +3645,10 @@ class WanVideoSampler: for pred_id, state in states.items(): name = state_names.get(pred_id, f"prediction_{pred_id}") if 'skipped_steps' in state: - log.info(f"TeaCache skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}") + log.info(f"{cache_type} skipped: {len(state['skipped_steps'])} {name} steps: {state['skipped_steps']}") transformer.teacache_state.clear_all() - - # if transformer.attention_mode == "spargeattn_tune": - # saved_state_dict = extract_sparse_attention_state_dict(transformer) - # torch.save(saved_state_dict, "sparge_wan.pt") - # save_torch_file(saved_state_dict, "sparge_wan.safetensors") + transformer.magcache_state.clear_all() + del states if force_offload: if model["manual_offloading"]: @@ -3617,7 +3670,7 @@ class WindowTracker: def __init__(self, verbose=False): self.window_map = {} # Maps frame sequence to persistent ID self.next_id = 0 - self.teacache_states = {} # Maps persistent ID to teacache state + self.cache_states = {} # Maps persistent ID to teacache state self.verbose = verbose def get_window_id(self, frames): @@ -3630,11 +3683,11 @@ class WindowTracker: return self.window_map[key] def get_teacache(self, window_id, base_state): - if window_id not in self.teacache_states: + if window_id not in self.cache_states: if self.verbose: log.info(f"Initializing persistent teacache for window {window_id}") - self.teacache_states[window_id] = base_state.copy() - return self.teacache_states[window_id] + self.cache_states[window_id] = base_state.copy() + return self.cache_states[window_id] #region VideoDecode class WanVideoDecode: @@ -3828,6 +3881,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoEnhanceAVideo": WanVideoEnhanceAVideo, "WanVideoContextOptions": WanVideoContextOptions, "WanVideoTeaCache": WanVideoTeaCache, + "WanVideoMagCache": WanVideoMagCache, "WanVideoVRAMManagement": WanVideoVRAMManagement, "WanVideoTextEmbedBridge": WanVideoTextEmbedBridge, "WanVideoFlowEdit": WanVideoFlowEdit, @@ -3868,6 +3922,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoEnhanceAVideo": "WanVideo Enhance-A-Video", "WanVideoContextOptions": "WanVideo Context Options", "WanVideoTeaCache": "WanVideo TeaCache", + "WanVideoMagCache": "WanVideo MagCache", "WanVideoVRAMManagement": "WanVideo VRAM Management", "WanVideoTextEmbedBridge": "WanVideo TextEmbed Bridge", "WanVideoFlowEdit": "WanVideo FlowEdit", diff --git a/skyreels/nodes.py b/skyreels/nodes.py index be49a94..663d37a 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -121,7 +121,7 @@ class WanVideoDiffusionForcingSampler: "samples": ("LATENT", {"tooltip": "init Latents to use for video2video process"} ), "prefix_samples": ("LATENT", {"tooltip": "prefix latents"} ), "denoise_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - "teacache_args": ("TEACACHEARGS", ), + "cache_args": ("CACHEARGS", ), "slg_args": ("SLGARGS", ), "rope_function": (["default", "comfy"], {"default": "comfy", "tooltip": "Comfy's RoPE implementation doesn't use complex numbers and can thus be compiled, that should be a lot faster when using torch.compile"}), "experimental_args": ("EXPERIMENTALARGS", ), @@ -135,7 +135,7 @@ class WanVideoDiffusionForcingSampler: CATEGORY = "WanVideoWrapper" def process(self, model, text_embeds, image_embeds, shift, fps, steps, addnoise_condition, cfg, seed, scheduler, - force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, slg_args=None, rope_function="default", teacache_args=None, + force_offload=True, samples=None, prefix_samples=None, denoise_strength=1.0, slg_args=None, rope_function="default", cache_args=None, teacache_args=None, experimental_args=None, unianimate_poses=None): #assert not (context_options and teacache_args), "Context options cannot currently be used together with teacache." patcher = model @@ -373,19 +373,29 @@ class WanVideoDiffusionForcingSampler: elif model["manual_offloading"]: transformer.to(device) - # Initialize TeaCache if enabled - if teacache_args is not None: - transformer.enable_teacache = True - 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"] - transformer.teacache_state.clear_all() - else: - transformer.enable_teacache = False + # Initialize Cache if enabled + transformer.enable_teacache = transformer.enable_magcache = False + if teacache_args is not None: #for backward compatibility on old workflows + cache_args = teacache_args + if cache_args is not None: + transformer.cache_device = cache_args["cache_device"] + if cache_args["cache_type"] == "TeaCache": + log.info(f"TeaCache: Using cache device: {transformer.cache_device}") + transformer.teacache_state.clear_all() + transformer.enable_teacache = True + transformer.rel_l1_thresh = cache_args["rel_l1_thresh"] + transformer.teacache_start_step = cache_args["start_step"] + transformer.teacache_end_step = len(init_timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.teacache_use_coefficients = cache_args["use_coefficients"] + transformer.teacache_mode = cache_args["mode"] + elif cache_args["cache_type"] == "MagCache": + log.info(f"MagCache: Using cache device: {transformer.cache_device}") + transformer.magcache_state.clear_all() + transformer.enable_magcache = True + transformer.magcache_start_step = cache_args["start_step"] + transformer.magcache_end_step = len(init_timesteps)-1 if cache_args["end_step"] == -1 else cache_args["end_step"] + transformer.magcache_thresh = cache_args["magcache_thresh"] + transformer.magcache_K = cache_args["magcache_K"] if slg_args is not None: transformer.slg_blocks = slg_args["blocks"] @@ -440,6 +450,8 @@ class WanVideoDiffusionForcingSampler: 'vace_data': vace_data, 'unianim_data': unianim_data, 'fps_embeds': fps_embeds, + "nag_params": text_embeds.get("nag_params", {}), + "nag_context": text_embeds.get("nag_prompt_embeds", None), } batch_size = 1 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index e3b3753..f48a952 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -404,12 +404,12 @@ class WanT2VCrossAttention(WanSelfAttention): self.attention_mode = attention_mode def forward(self, x, context, context_lens, clip_embed=None, audio_proj=None, audio_context_lens=None, audio_scale=1.0, - num_latent_frames=21, nag_params={}, nag_context=None): + num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False): b, n, d = x.size(0), self.num_heads, self.head_dim # compute query q = self.norm_q(self.q(x)).view(b, -1, n, d) - if nag_context is not None: + if nag_context is not None and not is_uncond: x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) else: k = self.norm_k(self.k(context)).view(b, -1, n, d) @@ -450,7 +450,7 @@ class WanI2VCrossAttention(WanSelfAttention): self.attention_mode = attention_mode def forward(self, x, context, context_lens, clip_embed, audio_proj=None, audio_context_lens=None, - audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None): + audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False): r""" Args: x(Tensor): Shape [B, L1, C] @@ -461,7 +461,7 @@ class WanI2VCrossAttention(WanSelfAttention): # compute query q = self.norm_q(self.q(x)).view(b, -1, n, d) - if nag_context is not None: + if nag_context is not None and not is_uncond: x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) else: # text attention @@ -583,7 +583,8 @@ class WanAttentionBlock(nn.Module): num_latent_frames=21, block_mask=None, nag_params={}, - nag_context=None + nag_context=None, + is_uncond=False ): r""" Args: @@ -640,15 +641,16 @@ class WanAttentionBlock(nn.Module): else: x = self.cross_attn_ffn(x, context, context_lens, e, clip_embed=clip_embed, grid_sizes=grid_sizes, audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale, - num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context) + num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond) del e return x #@torch.compiler.disable() def cross_attn_ffn(self, x, context, context_lens, e, clip_embed=None, grid_sizes=None, - audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None): + audio_proj=None, audio_context_lens=None, audio_scale=1.0, num_latent_frames=21, nag_params={}, + nag_context=None, is_uncond=False): x = x + self.cross_attn(self.norm3(x), context, context_lens, clip_embed=clip_embed, audio_proj=audio_proj, audio_context_lens=audio_context_lens, audio_scale=audio_scale, - num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context) + num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond) y = self.ffn(self.norm2(x) * (1 + e[4]) + e[3]) x = x + (y * e[5]) return x @@ -866,6 +868,7 @@ class WanModel(ModelMixin, ConfigMixin): main_device=torch.device('cuda'), offload_device=torch.device('cpu'), teacache_coefficients=[], + magcache_ratios=[], vace_layers=None, vace_in_dim=None, inject_sample_info=False, @@ -937,17 +940,27 @@ class WanModel(ModelMixin, ConfigMixin): self.offload_img_emb = False self.vace_blocks_to_swap = -1 + self.cache_device = offload_device + #init TeaCache variables self.enable_teacache = False self.rel_l1_thresh = 0.15 self.teacache_start_step= 0 self.teacache_end_step = -1 - self.teacache_cache_device = offload_device - self.teacache_state = TeaCacheState(cache_device=self.teacache_cache_device) + self.teacache_state = TeaCacheState(cache_device=self.cache_device) self.teacache_coefficients = teacache_coefficients self.teacache_use_coefficients = False self.teacache_mode = 'e' + #init MagCache variables + self.enable_magcache = False + self.magcache_state = MagCacheState(cache_device=self.cache_device) + self.magcache_thresh = 0.24 + self.magcache_K = 4 + self.magcache_start_step = 0 + self.magcache_end_step = -1 + self.magcache_ratios = magcache_ratios + self.slg_blocks = None self.slg_start_percent = 0.0 self.slg_end_percent = 1.0 @@ -1175,11 +1188,12 @@ class WanModel(ModelMixin, ConfigMixin): seq_len, is_uncond=False, current_step_percentage=0.0, + current_step=0, + total_steps=50, clip_fea=None, y=None, device=torch.device('cuda'), freqs=None, - current_step=0, pred_id=None, control_lora_enabled=False, vace_data=None, @@ -1407,9 +1421,7 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) if self.enable_teacache and self.teacache_start_step <= current_step <= self.teacache_end_step: if pred_id is None: - pred_id = self.teacache_state.new_prediction(cache_device=self.teacache_cache_device) - #log.info(current_step) - #log.info(f"TeaCache: Initializing TeaCache variables for model pred: {pred_id}") + pred_id = self.teacache_state.new_prediction(cache_device=self.cache_device) should_calc = True else: previous_modulated_input = self.teacache_state.get(pred_id)['previous_modulated_input'] @@ -1429,29 +1441,64 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(e0.device) + temb_relative_l1 del temb_relative_l1 - #print("accumulated_rel_l1_distance", accumulated_rel_l1_distance) if accumulated_rel_l1_distance < self.rel_l1_thresh: should_calc = False 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) + accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.cache_device) - previous_modulated_input = e.to(self.teacache_cache_device).clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.to(self.teacache_cache_device).clone() + previous_modulated_input = e.to(self.cache_device).clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.to(self.cache_device).clone() 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}") self.teacache_state.update( pred_id, accumulated_rel_l1_distance=accumulated_rel_l1_distance, ) self.teacache_state.get(pred_id)['skipped_steps'].append(current_step) - if not self.enable_teacache or (self.enable_teacache and should_calc): - if self.enable_teacache: - original_x = x.to(self.teacache_cache_device).clone() + # enable magcache + if self.enable_magcache and self.magcache_start_step <= current_step <= self.magcache_end_step: + if pred_id is None: + pred_id = self.magcache_state.new_prediction(cache_device=self.cache_device) + should_calc = True + else: + accumulated_ratio = self.magcache_state.get(pred_id)['accumulated_ratio'] + accumulated_err = self.magcache_state.get(pred_id)['accumulated_err'] + accumulated_steps = self.magcache_state.get(pred_id)['accumulated_steps'] + + calibration_len = len(self.magcache_ratios) // 2 + cur_mag_ratio = self.magcache_ratios[int((current_step*(calibration_len/total_steps)))] + + accumulated_ratio *= cur_mag_ratio + accumulated_err += np.abs(1-accumulated_ratio) + accumulated_steps += 1 + + self.magcache_state.update( + pred_id, + accumulated_ratio=accumulated_ratio, + accumulated_steps=accumulated_steps, + accumulated_err=accumulated_err + ) + + if accumulated_err<=self.magcache_thresh and accumulated_steps<=self.magcache_K: + should_calc = False + x += self.magcache_state.get(pred_id)['residual_cache'].to(x.device) + self.magcache_state.get(pred_id)['skipped_steps'].append(current_step) + else: + should_calc = True + self.magcache_state.update( + pred_id, + accumulated_ratio=1.0, + accumulated_steps=0, + accumulated_err=0 + ) + + if should_calc: + if self.enable_teacache or self.enable_magcache: + original_x = x.to(self.cache_device).clone() if hasattr(self, "dwpose_embedding") and unianim_data is not None: if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']: @@ -1476,7 +1523,8 @@ class WanModel(ModelMixin, ConfigMixin): audio_scale=audio_scale, block_mask=self.block_mask, nag_params=nag_params, - nag_context=nag_context + nag_context=nag_context, + is_uncond = is_uncond ) if vace_data is not None: @@ -1539,7 +1587,12 @@ class WanModel(ModelMixin, ConfigMixin): accumulated_rel_l1_distance=accumulated_rel_l1_distance, previous_modulated_input=previous_modulated_input ) - + elif self.enable_magcache and (self.magcache_start_step <= current_step <= self.magcache_end_step) and pred_id is not None: + self.magcache_state.update( + pred_id, + residual_cache=(x.to(original_x.device) - original_x) + ) + 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:] @@ -1607,14 +1660,40 @@ class TeaCacheState: def get(self, pred_id): return self.states.get(pred_id, {}) - - def report(self): - for pred_id in self.states: - log.info(f"Prediction {pred_id}: {self.states[pred_id]}") - def clear_prediction(self, pred_id): - if pred_id in self.states: - del self.states[pred_id] + def clear_all(self): + self.states = {} + self._next_pred_id = 0 + +class MagCacheState: + def __init__(self, cache_device='cpu'): + self.cache_device = cache_device + self.states = {} + self._next_pred_id = 0 + + def new_prediction(self, cache_device='cpu'): + """Create new prediction state and return its ID""" + self.cache_device = cache_device + pred_id = self._next_pred_id + self._next_pred_id += 1 + self.states[pred_id] = { + 'residual_cache': None, + 'accumulated_ratio': 1.0, + 'accumulated_steps': 0, + 'accumulated_err': 0, + 'skipped_steps': [], + } + return pred_id + + def update(self, pred_id, **kwargs): + """Update state for specific prediction""" + if pred_id not in self.states: + return None + for key, value in kwargs.items(): + self.states[pred_id][key] = value + + def get(self, pred_id): + return self.states.get(pred_id, {}) def clear_all(self): self.states = {}