Make uni3c offloading optional
This commit is contained in:
+3
-3
@@ -334,7 +334,7 @@ class WanControlNet(ModelMixin):
|
||||
|
||||
self.controlnet_mask_embedding = MaskCamEmbed(controlnet_cfg)
|
||||
|
||||
def forward(self, render_latent, render_mask, camera_embedding, temb, device):
|
||||
def forward(self, render_latent, render_mask, camera_embedding, temb, out_device):
|
||||
controlnet_rotary_emb = self.controlnet_rope(render_latent)
|
||||
controlnet_inputs = self.controlnet_patch_embedding(render_latent.to(torch.float32))
|
||||
if not self.quantized:
|
||||
@@ -354,7 +354,7 @@ class WanControlNet(ModelMixin):
|
||||
if add_inputs is not None:
|
||||
add_inputs = self.controlnet_mask_embedding(add_inputs)
|
||||
controlnet_inputs = controlnet_inputs + add_inputs
|
||||
|
||||
|
||||
hidden_states = self.proj_in(controlnet_inputs)
|
||||
|
||||
controlnet_states = []
|
||||
@@ -364,6 +364,6 @@ class WanControlNet(ModelMixin):
|
||||
temb=temb,
|
||||
rotary_emb=controlnet_rotary_emb
|
||||
)
|
||||
controlnet_states.append(self.proj_out[i](hidden_states).to(device))
|
||||
controlnet_states.append(self.proj_out[i](hidden_states).to(out_device))
|
||||
|
||||
return controlnet_states
|
||||
|
||||
+29
-33
@@ -10,8 +10,6 @@ from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
import folder_paths
|
||||
|
||||
import json
|
||||
import numpy as np
|
||||
|
||||
class WanVideoUni3C_ControlnetLoader:
|
||||
@classmethod
|
||||
@@ -22,7 +20,7 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
|
||||
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
|
||||
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "offload_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
"attention_mode": ([
|
||||
"sdpa",
|
||||
"sageattn",
|
||||
@@ -45,17 +43,17 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
transformer_load_device = device if load_device == "main_device" else offload_device
|
||||
|
||||
|
||||
base_dtype = {"fp8_e4m3fn": torch.float8_e4m3fn, "fp8_e4m3fn_fast": torch.float8_e4m3fn, "bf16": torch.bfloat16, "fp16": torch.float16, "fp16_fast": torch.float16, "fp32": torch.float32}[base_precision]
|
||||
|
||||
|
||||
|
||||
model_path = folder_paths.get_full_path_or_raise("controlnet", model)
|
||||
|
||||
|
||||
sd = load_torch_file(model_path, device=transformer_load_device, safe_load=True)
|
||||
|
||||
if not "controlnet_patch_embedding.weight" in sd:
|
||||
raise ValueError("Invalid ControlNet model")
|
||||
|
||||
|
||||
in_channels = sd["controlnet_patch_embedding.weight"].shape[1]
|
||||
ffn_dim = sd["controlnet_blocks.0.ffn.0.bias"].shape[0]
|
||||
|
||||
@@ -79,7 +77,7 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
with init_empty_weights():
|
||||
controlnet = WanControlNet(controlnet_cfg)
|
||||
controlnet.eval()
|
||||
|
||||
|
||||
if quantization == "disabled":
|
||||
for k, v in sd.items():
|
||||
if isinstance(v, torch.Tensor):
|
||||
@@ -97,18 +95,18 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "head", "time_in", "vector_in", "controlnet_patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "proj_in"}
|
||||
|
||||
|
||||
log.info("Using accelerate to load and assign controlnet model weights to device...")
|
||||
param_count = sum(1 for _ in controlnet.named_parameters())
|
||||
for name, param in tqdm(controlnet.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
for name, param in tqdm(controlnet.named_parameters(),
|
||||
desc=f"Loading transformer parameters to {transformer_load_device}",
|
||||
total=param_count,
|
||||
leave=True):
|
||||
dtype_to_use = base_dtype if any(keyword in name for keyword in params_to_keep) else dtype
|
||||
if "controlnet_patch_embedding" in name:
|
||||
dtype_to_use = torch.float32
|
||||
set_module_tensor_to_device(controlnet, name, device=transformer_load_device, dtype=dtype_to_use, value=sd[name])
|
||||
|
||||
|
||||
del sd
|
||||
|
||||
if compile_args is not None:
|
||||
@@ -123,8 +121,8 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
for i, block in enumerate(controlnet.controlnet_blocks):
|
||||
controlnet.controlnet_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
else:
|
||||
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
controlnet = torch.compile(controlnet, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
|
||||
|
||||
if load_device == "offload_device" and controlnet.device != offload_device:
|
||||
log.info(f"Moving controlnet model from {controlnet.device} to {offload_device}")
|
||||
@@ -146,6 +144,7 @@ class WanVideoUni3C_embeds:
|
||||
"optional": {
|
||||
"render_latent": ("LATENT",),
|
||||
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
|
||||
"offload": ("BOOLEAN", {"default": True, "tooltip": "If enabled, the controlnet model will be offloaded before main model block processing to save VRAM."}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -154,9 +153,7 @@ class WanVideoUni3C_embeds:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None, offload=True):
|
||||
|
||||
latent_mask = latents = None
|
||||
if render_latent is not None:
|
||||
@@ -164,17 +161,17 @@ class WanVideoUni3C_embeds:
|
||||
# nframe = latents.shape[2] * 4
|
||||
# height = latents.shape[3] * 8
|
||||
# width = latents.shape[4] * 8
|
||||
|
||||
|
||||
if render_mask is not None:
|
||||
raise NotImplementedError("render_mask is not implemented at this time")
|
||||
mask = torch.nn.functional.interpolate(
|
||||
render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(nframe, height, width),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).squeeze(0)
|
||||
latent_mask = mask.unsqueeze(0).to(device)
|
||||
log.info(f"latent mask shape {latent_mask.shape}")
|
||||
# mask = torch.nn.functional.interpolate(
|
||||
# render_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
# size=(nframe, height, width),
|
||||
# mode='trilinear',
|
||||
# align_corners=False
|
||||
# ).squeeze(0)
|
||||
# latent_mask = mask.unsqueeze(0).to(device)
|
||||
# log.info(f"latent mask shape {latent_mask.shape}")
|
||||
|
||||
# # load camera
|
||||
# cam_info = json.load(open(f"{render_path}/cam_info.json"))
|
||||
@@ -199,7 +196,7 @@ class WanVideoUni3C_embeds:
|
||||
# K_inv = K.inverse()
|
||||
# intrinsic = K[None].repeat(nframe, 1, 1)
|
||||
|
||||
|
||||
|
||||
# w2c_0, c2w_0 = set_initial_camera(start_elevation, depth_avg)
|
||||
# w2cs, c2ws, intrinsic = build_cameras(cam_traj=cam_traj,
|
||||
# w2c_0=w2c_0,
|
||||
@@ -215,7 +212,7 @@ class WanVideoUni3C_embeds:
|
||||
# y_offset=y_offset,
|
||||
# z_offset=z_offset)
|
||||
|
||||
|
||||
|
||||
# from .camera import get_camera_embedding
|
||||
# camera_embedding = get_camera_embedding(intrinsic, w2cs, nframe, height, width, normalize=True)
|
||||
#print("camera embedding shape", camera_embedding.shape)
|
||||
@@ -227,11 +224,12 @@ class WanVideoUni3C_embeds:
|
||||
"end": end_percent,
|
||||
"render_latent": latents,
|
||||
"render_mask": latent_mask,
|
||||
"camera_embedding": None
|
||||
"camera_embedding": None,
|
||||
"offload": offload,
|
||||
}
|
||||
|
||||
|
||||
return (uni3c_embeds,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoUni3C_ControlnetLoader": WanVideoUni3C_ControlnetLoader,
|
||||
"WanVideoUni3C_embeds": WanVideoUni3C_embeds,
|
||||
@@ -240,5 +238,3 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoUni3C_ControlnetLoader": "WanVideo Uni3C Controlnet Loader",
|
||||
"WanVideoUni3C_embeds": "WanVideo Uni3C Embeds",
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user