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
This commit is contained in:
kijai
2025-06-13 12:31:13 +03:00
parent e7c39757f6
commit ad43eed0d0
20 changed files with 661 additions and 503 deletions
@@ -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
]
@@ -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
},
{
@@ -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
},
{
@@ -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
},
{
@@ -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
},
{
@@ -907,9 +907,9 @@
"link": null
},
{
"name": "teacache_args",
"name": "cache_args",
"shape": 7,
"type": "TEACACHEARGS",
"type": "CACHEARGS",
"link": null
},
{
@@ -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
},
{
@@ -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
},
{
@@ -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
},
{
@@ -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
]
@@ -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
},
{
@@ -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
},
{
@@ -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
},
{
@@ -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
},
{
@@ -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
},
{
@@ -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",
@@ -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
},
{
+132 -77
View File
@@ -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",
+27 -15
View File
@@ -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
+109 -30
View File
@@ -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 = {}