Make I2V resolution adjustment optional and other tweaks

This commit is contained in:
kijai
2025-03-01 16:58:10 +02:00
parent 66c8ebe4d3
commit df6d0b677f
4 changed files with 443 additions and 315 deletions
@@ -1,78 +1,7 @@
{
"last_node_id": 36,
"last_link_id": 40,
"last_node_id": 38,
"last_link_id": 42,
"nodes": [
{
"id": 21,
"type": "WanVideoVAELoader",
"pos": [
401.8250427246094,
393.2132873535156
],
"size": [
315,
82
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "VAE",
"links": [
21,
34
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoVAELoader"
},
"widgets_values": [
"wanvideo\\Wan2_1_VAE_bf16.safetensors",
"bf16"
]
},
{
"id": 18,
"type": "LoadImage",
"pos": [
417.39227294921875,
529.9345092773438
],
"size": [
315,
314
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
18
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"oldman_upscaled.png",
"image"
]
},
{
"id": 32,
"type": "WanVideoBlockSwap",
@@ -82,10 +11,10 @@
],
"size": [
315,
58
106
],
"flags": {},
"order": 2,
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
@@ -102,7 +31,9 @@
"Node name for S&R": "WanVideoBlockSwap"
},
"widgets_values": [
10
10,
false,
false
]
},
{
@@ -114,10 +45,10 @@
],
"size": [
377.1661376953125,
115.05413055419922
130
],
"flags": {},
"order": 3,
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
@@ -136,7 +67,8 @@
"widgets_values": [
"umt5-xxl-enc-bf16.safetensors",
"bf16",
"offload_device"
"offload_device",
"disabled"
]
},
{
@@ -191,7 +123,7 @@
106
],
"flags": {},
"order": 4,
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
@@ -214,62 +146,293 @@
]
},
{
"id": 17,
"type": "WanVideoImageClipEncode",
"id": 35,
"type": "WanVideoTorchCompileSettings",
"pos": [
875.01025390625,
278.4588623046875
124.67726135253906,
-627.7935180664062
],
"size": [
315,
170
390.5999755859375,
178
],
"flags": {},
"order": 11,
"order": 3,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "torch_compile_args",
"type": "WANCOMPILEARGS",
"links": [],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoTorchCompileSettings"
},
"widgets_values": [
"inductor",
false,
"default",
false,
64,
true
]
},
{
"id": 37,
"type": "Note",
"pos": [
170.94039916992188,
-750.3395385742188
],
"size": [
303.0501403808594,
65.95614624023438
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"If you have Triton installed, connect this for ~30% speed increase"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 36,
"type": "Note",
"pos": [
712.7784423828125,
-580.95947265625
],
"size": [
378.52294921875,
144.30191040039062
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"fp8_fast seems to cause huge quality degradation\n\nfp_16_fast enables \"Full FP16 Accmumulation in FP16 GEMMs\" feature available in the very latest pytorch nightly, this is around 20% speed boost. \n\nSageattn if you have it installed can be used for almost double inference speed"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 22,
"type": "WanVideoModelLoader",
"pos": [
620.3950805664062,
-357.8426818847656
],
"size": [
477.4410095214844,
226.43276977539062
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "clip",
"type": "WANCLIP",
"link": 17
"name": "compile_args",
"type": "WANCOMPILEARGS",
"shape": 7,
"link": null
},
{
"name": "image",
"type": "IMAGE",
"link": 18
"name": "block_swap_args",
"type": "BLOCKSWAPARGS",
"shape": 7,
"link": 39
},
{
"name": "vae",
"type": "VAE",
"link": 21
"name": "lora",
"type": "WANVIDLORA",
"shape": 7,
"link": null
}
],
"outputs": [
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"name": "model",
"type": "WANVIDEOMODEL",
"links": [
32
29
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoImageClipEncode"
"Node name for S&R": "WanVideoModelLoader"
},
"widgets_values": [
512,
512,
81,
true
"WanVideo\\Wan2_1-I2V-14B-480P_fp8_e4m3fn.safetensors",
"fp16",
"fp8_e4m3fn",
"offload_device",
"sdpa"
]
},
{
"id": 27,
"type": "WanVideoSampler",
"pos": [
1349.4732666015625,
-233.91836547851562
],
"size": [
315,
534.1923217773438
],
"flags": {},
"order": 12,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 29
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 30
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 32
},
{
"name": "samples",
"type": "LATENT",
"shape": 7,
"link": null
},
{
"name": "feta_args",
"type": "FETAARGS",
"shape": 7,
"link": null
},
{
"name": "context_options",
"type": "WANVIDCONTEXT",
"shape": 7,
"link": null
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"links": [
33
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoSampler"
},
"widgets_values": [
10,
6,
5,
123,
"fixed",
true,
"unipc",
0,
1,
""
]
},
{
"id": 21,
"type": "WanVideoVAELoader",
"pos": [
361.2858581542969,
388.70892333984375
],
"size": [
315,
82
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "vae",
"type": "WANVAE",
"links": [
21,
34
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoVAELoader"
},
"widgets_values": [
"wanvideo\\Wan2_1_VAE_bf16.safetensors",
"bf16"
]
},
{
"id": 18,
"type": "LoadImage",
"pos": [
367.8443298339844,
527.23193359375
],
"size": [
315,
314
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "IMAGE",
"type": "IMAGE",
"links": [
18
]
},
{
"name": "MASK",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "LoadImage"
},
"widgets_values": [
"oldman_upscaled.png",
"image"
]
},
{
"id": 28,
"type": "WanVideoDecode",
"pos": [
1319.96875,
11.251319885253906
1346.09375,
-467.1110534667969
],
"size": [
315,
@@ -281,7 +444,7 @@
"inputs": [
{
"name": "vae",
"type": "VAE",
"type": "WANVAE",
"link": 34
},
{
@@ -295,7 +458,7 @@
"name": "images",
"type": "IMAGE",
"links": [
36
41
],
"slot_index": 0
}
@@ -312,47 +475,74 @@
]
},
{
"id": 33,
"type": "Note",
"id": 38,
"type": "GetImageSizeAndCount",
"pos": [
227.3764190673828,
-205.28524780273438
1703.7633056640625,
-346.5655517578125
],
"size": [
351.70458984375,
60
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 30,
"type": "VHS_VideoCombine",
"pos": [
1712.420654296875,
-353.76507568359375
],
"size": [
214.7587890625,
542.7587890625
277.20001220703125,
86
],
"flags": {},
"order": 14,
"mode": 0,
"inputs": [
{
"name": "image",
"type": "IMAGE",
"link": 41
}
],
"outputs": [
{
"name": "image",
"type": "IMAGE",
"links": [
42
],
"slot_index": 0
},
{
"name": "width",
"type": "INT",
"links": null
},
{
"name": "height",
"type": "INT",
"links": null
},
{
"name": "count",
"type": "INT",
"links": null
}
],
"properties": {
"Node name for S&R": "GetImageSizeAndCount"
}
},
{
"id": 30,
"type": "VHS_VideoCombine",
"pos": [
1988.0869140625,
-413.2225341796875
],
"size": [
557.9904174804688,
885.9904174804688
],
"flags": {},
"order": 15,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 36
"link": 42
},
{
"name": "audio",
@@ -398,194 +588,82 @@
"hidden": false,
"paused": false,
"params": {
"filename": "WanVideo2_1_00011.mp4",
"filename": "WanVideo2_1_00115.mp4",
"subfolder": "",
"type": "output",
"format": "video/h264-mp4",
"frame_rate": 16,
"workflow": "WanVideo2_1_00011.png",
"fullpath": "N:\\AI\\ComfyUI\\output\\WanVideo2_1_00011.mp4"
"workflow": "WanVideo2_1_00115.png",
"fullpath": "N:\\AI\\ComfyUI\\output\\WanVideo2_1_00115.mp4"
}
}
}
},
{
"id": 27,
"type": "WanVideoSampler",
"id": 17,
"type": "WanVideoImageClipEncode",
"pos": [
1315.2401123046875,
-356.4367980957031
875.01025390625,
278.4588623046875
],
"size": [
315,
242
266
],
"flags": {},
"order": 12,
"order": 11,
"mode": 0,
"inputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"link": 29
"name": "clip",
"type": "WANCLIP",
"link": 17
},
{
"name": "text_embeds",
"type": "WANVIDEOTEXTEMBEDS",
"link": 30
"name": "image",
"type": "IMAGE",
"link": 18
},
{
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"link": 32
"name": "vae",
"type": "WANVAE",
"link": 21
}
],
"outputs": [
{
"name": "samples",
"type": "LATENT",
"name": "image_embeds",
"type": "WANVIDIMAGE_EMBEDS",
"links": [
33
32
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoSampler"
"Node name for S&R": "WanVideoImageClipEncode"
},
"widgets_values": [
10,
6,
5,
1057359483639286,
"fixed",
832,
480,
81,
true,
"unipc"
]
},
{
"id": 35,
"type": "WanVideoTorchCompileSettings",
"pos": [
124.67726135253906,
-627.7935180664062
],
"size": [
390.5999755859375,
178
],
"flags": {},
"order": 6,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "torch_compile_args",
"type": "COMPILEARGS",
"links": [],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoTorchCompileSettings"
},
"widgets_values": [
"inductor",
false,
"default",
false,
64,
0,
1,
1,
true
]
},
{
"id": 22,
"type": "WanVideoModelLoader",
"pos": [
620.3950805664062,
-357.8426818847656
],
"size": [
477.4410095214844,
226.43276977539062
],
"flags": {},
"order": 9,
"mode": 0,
"inputs": [
{
"name": "compile_args",
"type": "COMPILEARGS",
"shape": 7,
"link": null
},
{
"name": "block_swap_args",
"type": "BLOCKSWAPARGS",
"shape": 7,
"link": 39
},
{
"name": "lora",
"type": "HYVIDLORA",
"shape": 7,
"link": null
}
],
"outputs": [
{
"name": "model",
"type": "WANVIDEOMODEL",
"links": [
29
],
"slot_index": 0
}
],
"properties": {
"Node name for S&R": "WanVideoModelLoader"
},
"widgets_values": [
"WanVideo\\Wan2_1-I2V-14B-480P_fp8_e4m3fn.safetensors",
"bf16",
"fp8_e4m3fn",
"offload_device",
"sageattn"
]
},
{
"id": 34,
"id": 33,
"type": "Note",
"pos": [
912.0381469726562,
501.89813232421875
234.58334350585938,
-189.06956481933594
],
"size": [
262.5184020996094,
58
],
"flags": {},
"order": 7,
"mode": 0,
"inputs": [],
"outputs": [],
"properties": {},
"widgets_values": [
"Under 81 frames doesn't seem to work?"
],
"color": "#432",
"bgcolor": "#653"
},
{
"id": 36,
"type": "Note",
"pos": [
796.0189208984375,
-521.5020751953125
],
"size": [
298.2554016113281,
108.62744140625
351.70458984375,
60
],
"flags": {},
"order": 8,
@@ -594,7 +672,7 @@
"outputs": [],
"properties": {},
"widgets_values": [
"sdpa should work too, haven't tested flaash\n\nfp8_fast seems to cause huge quality degradation"
"Models:\nhttps://huggingface.co/Kijai/WanVideo_comfy/tree/main"
],
"color": "#432",
"bgcolor": "#653"
@@ -673,14 +751,6 @@
0,
"VAE"
],
[
36,
28,
0,
30,
0,
"IMAGE"
],
[
39,
32,
@@ -688,6 +758,22 @@
22,
1,
"BLOCKSWAPARGS"
],
[
41,
28,
0,
38,
0,
"IMAGE"
],
[
42,
38,
0,
30,
0,
"IMAGE"
]
],
"groups": [],
@@ -696,13 +782,14 @@
"ds": {
"scale": 0.672749994932598,
"offset": [
427.5980150909793,
736.9619463805877
341.6634266370342,
770.6955320763833
]
},
"node_versions": {
"ComfyUI-WanVideoWrapper": "c83f47e4d97b5891058555df16db5e33d16afab1",
"comfy-core": "0.3.14",
"ComfyUI-WanVideoWrapper": "66c8ebe4d375e5dc1a510745931ae9bd14ec7887",
"comfy-core": "0.3.18",
"ComfyUI-KJNodes": "9a15e22f5e9416c0968ce3de33923f8f601257dd",
"ComfyUI-VideoHelperSuite": "2c25b8b53835aaeb63f831b3137c705cf9f85dce"
},
"VHS_latentpreview": true,
+46 -17
View File
@@ -12,6 +12,7 @@ from .wanvideo.modules.t5 import T5EncoderModel
from .wanvideo.utils.fm_solvers import (FlowDPMSolverMultistepScheduler,
get_sampling_sigmas, retrieve_timesteps)
from .wanvideo.utils.fm_solvers_unipc import FlowUniPCMultistepScheduler
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
from .enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight, set_num_frames
@@ -39,6 +40,8 @@ class WanVideoBlockSwap:
return {
"required": {
"blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 40, "step": 1, "tooltip": "Number of double blocks to swap"}),
"offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}),
"offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}),
},
}
RETURN_TYPES = ("BLOCKSWAPARGS",)
@@ -708,6 +711,7 @@ class WanVideoImageClipEncode:
"noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation, helpful for I2V where some noise can add motion and give sharper results"}),
"latent_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional latent multiplier, helpful for I2V where lower values allow for more motion"}),
"clip_embed_strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Additional clip embed multiplier"}),
"adjust_resolution": ("BOOLEAN", {"default": True, "tooltip": "Performs the same resolution adjustment as in the original code"}),
}
}
@@ -717,7 +721,8 @@ class WanVideoImageClipEncode:
FUNCTION = "process"
CATEGORY = "WanVideoWrapper"
def process(self, clip, vae, image, num_frames, generation_width, generation_height, force_offload=True, noise_aug_strength=0.0, latent_strength=1.0, clip_embed_strength=1.0):
def process(self, clip, vae, image, num_frames, generation_width, generation_height, force_offload=True, noise_aug_strength=0.0,
latent_strength=1.0, clip_embed_strength=1.0, adjust_resolution=True):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -742,15 +747,21 @@ class WanVideoImageClipEncode:
clip.model.to(offload_device)
mm.soft_empty_cache()
aspect_ratio = H / W
lat_h = round(
if adjust_resolution:
aspect_ratio = H / W
lat_h = round(
np.sqrt(max_area * aspect_ratio) // vae_stride[1] //
patch_size[1] * patch_size[1])
lat_w = round(
np.sqrt(max_area / aspect_ratio) // vae_stride[2] //
patch_size[2] * patch_size[2])
h = lat_h * vae_stride[1]
w = lat_w * vae_stride[2]
lat_w = round(
np.sqrt(max_area / aspect_ratio) // vae_stride[2] //
patch_size[2] * patch_size[2])
h = lat_h * vae_stride[1]
w = lat_w * vae_stride[2]
else:
h = generation_height
w = generation_width
lat_h = h // 8
lat_w = w // 8
# Step 1: Create initial mask with ones for first frame, zeros for others
mask = torch.ones(1, num_frames, lat_h, lat_w, device=device)
@@ -888,7 +899,7 @@ class WanVideoSampler:
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"force_offload": ("BOOLEAN", {"default": True}),
"scheduler": (["unipc", "dpm++", "dpm++_sde"],
"scheduler": (["unipc", "dpm++", "dpm++_sde", "euler"],
{
"default": 'dpm++'
}),
@@ -928,6 +939,16 @@ class WanVideoSampler:
sample_scheduler.set_timesteps(
steps, device=device, shift=shift)
timesteps = sample_scheduler.timesteps
elif scheduler == 'euler':
sample_scheduler = FlowMatchEulerDiscreteScheduler(
num_train_timesteps=1000,
shift=shift,
use_dynamic_shifting=False)
sampling_sigmas = get_sampling_sigmas(steps, shift)
timesteps, _ = retrieve_timesteps(
sample_scheduler,
device=device,
sigmas=sampling_sigmas)
elif 'dpm++' in scheduler:
if scheduler == 'dpm++_sde':
algorithm_type = "sde-dpmsolver++"
@@ -1060,9 +1081,15 @@ class WanVideoSampler:
for name, param in transformer.named_parameters():
if "block" not in name:
param.data = param.data.to(device)
elif model["block_swap_args"]["offload_txt_emb"] and "txt_emb" in name:
param.data = param.data.to(offload_device)
elif model["block_swap_args"]["offload_img_emb"] and "img_emb" in name:
param.data = param.data.to(offload_device)
transformer.block_swap(
model["block_swap_args"]["blocks_to_swap"] - 1 ,
model["block_swap_args"]["offload_txt_emb"],
model["block_swap_args"]["offload_img_emb"],
)
else:
if model["manual_offloading"]:
@@ -1108,6 +1135,8 @@ class WanVideoSampler:
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*8}x{latent.shape[2]*8} with {steps} steps")
intermediate_device = device
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=model["dtype"], enabled=True):
for i, t in enumerate(tqdm(timesteps)):
latent_model_input = [latent.to(device)]
@@ -1124,8 +1153,8 @@ class WanVideoSampler:
if context_options is not None:
counter = torch.zeros_like(latent_model_input[0], device=offload_device)
noise_pred = torch.zeros_like(latent_model_input[0], device=offload_device)
counter = torch.zeros_like(latent_model_input[0], device=intermediate_device)
noise_pred = torch.zeros_like(latent_model_input[0], device=intermediate_device)
context_queue = list(context(
i, steps, latent_video_length, context_frames, context_stride, context_overlap,
))
@@ -1134,10 +1163,10 @@ class WanVideoSampler:
partial_latent_model_input = [latent_model_input[0][:, c, :, :]]
# Model inference - returns [frames, channels, height, width]
noise_pred_cond = transformer(
partial_latent_model_input, t=timestep, **arg_c)[0].to(offload_device)
partial_latent_model_input, t=timestep, **arg_c)[0].to(intermediate_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
partial_latent_model_input, t=timestep, **arg_null)[0].to(offload_device)
partial_latent_model_input, t=timestep, **arg_null)[0].to(intermediate_device)
noise_pred_context = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
@@ -1164,10 +1193,10 @@ class WanVideoSampler:
else:
#model inference start
noise_pred_cond = transformer(
latent_model_input, t=timestep, **arg_c)[0].to(offload_device)
latent_model_input, t=timestep, **arg_c)[0].to(intermediate_device)
if cfg[i] != 1.0:
noise_pred_uncond = transformer(
latent_model_input, t=timestep, **arg_null)[0].to(offload_device)
latent_model_input, t=timestep, **arg_null)[0].to(intermediate_device)
noise_pred = noise_pred_uncond + cfg[i] * (
noise_pred_cond - noise_pred_uncond)
@@ -1175,7 +1204,7 @@ class WanVideoSampler:
noise_pred = noise_pred_cond
#model inference end
latent = latent.to(offload_device)
latent = latent.to(intermediate_device)
temp_x0 = sample_scheduler.step(
noise_pred.unsqueeze(0),
@@ -1188,7 +1217,7 @@ class WanVideoSampler:
x0 = [latent.to(device)]
if callback is not None:
callback_latent = (latent_model_input[0].cpu() - noise_pred * t.cpu() / 1000).detach().permute(1,0,2,3)
callback_latent = (latent_model_input[0] - noise_pred.to(t.device) * t / 1000).detach().permute(1,0,2,3)
callback(i, callback_latent, None, steps)
else:
pbar.update(1)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-WanVideoWrapper"
description = "ComfyUI diffusers wrapper nodes for WanVideo"
version = "1.0.1"
version = "1.0.2"
license = {file = "LICENSE"}
dependencies = ["accelerate >= 1.2.1", "diffusers >= 0.31.0", "ftfy"]
+13 -1
View File
@@ -488,6 +488,8 @@ class WanModel(ModelMixin, ConfigMixin):
self.offload_device = offload_device
self.blocks_to_swap = -1
self.offload_txt_emb = False
self.offload_img_emb = False
# embeddings
self.patch_embedding = nn.Conv3d(
@@ -522,9 +524,11 @@ class WanModel(ModelMixin, ConfigMixin):
# initialize weights
#self.init_weights()
def block_swap(self, blocks_to_swap):
def block_swap(self, blocks_to_swap, offload_txt_emb=False, offload_img_emb=False):
print(f"Swapping {blocks_to_swap + 1} transformer blocks")
self.blocks_to_swap = blocks_to_swap
self.offload_img_emb = offload_img_emb
self.offload_txt_emb = offload_txt_emb
for b, block in tqdm(enumerate(self.blocks), total=len(self.blocks), desc="Initializing block swap"):
if b > self.blocks_to_swap:
@@ -595,16 +599,24 @@ class WanModel(ModelMixin, ConfigMixin):
# context
context_lens = None
if self.offload_txt_emb:
self.text_embedding.to(self.main_device)
context = self.text_embedding(
torch.stack([
torch.cat(
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in context
]))
if self.offload_txt_emb:
self.text_embedding.to(self.offload_device, non_blocking=True)
if clip_fea is not None:
if self.offload_img_emb:
self.img_emb.to(self.main_device)
context_clip = self.img_emb(clip_fea) # bs x 257 x dim
context = torch.concat([context_clip, context], dim=1)
if self.offload_img_emb:
self.img_emb.to(self.offload_device, non_blocking=True)
# arguments
kwargs = dict(