From 3197107b5ea30c0ed26907afa19501fe4a514d92 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 4 Dec 2024 03:53:51 +0200 Subject: [PATCH] fixup block_swap start --- hyvideo/modules/models.py | 11 ++++++----- nodes.py | 10 +++++++++- 2 files changed, 15 insertions(+), 6 deletions(-) diff --git a/hyvideo/modules/models.py b/hyvideo/modules/models.py index 0fdaf5a..69bdde0 100644 --- a/hyvideo/modules/models.py +++ b/hyvideo/modules/models.py @@ -443,6 +443,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): text_states_dim_2: int = 768, dtype: Optional[torch.dtype] = None, device: Optional[torch.device] = None, + main_device: Optional[torch.device] = None, offload_device: Optional[torch.device] = None, attention_mode: str = "flash_attn", ): @@ -456,7 +457,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.guidance_embed = guidance_embed self.rope_dim_list = rope_dim_list - self.main_device = device + self.main_device = main_device self.offload_device = offload_device # Text projection. Default to linear projection. @@ -572,11 +573,11 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): self.single_blocks_to_swap = single_blocks_to_swap for b, block in enumerate(self.double_blocks): if b < 0 or b > self.double_blocks_to_swap: - mm.soft_empty_cache() + #mm.soft_empty_cache() block.to(self.main_device) for b, block in enumerate(self.single_blocks): if b < 0 or b > self.single_blocks_to_swap: - mm.soft_empty_cache() + #mm.soft_empty_cache() block.to(self.main_device) def enable_deterministic(self): @@ -669,7 +670,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): img, txt = block(*double_block_args) if b >= 0 and b <= self.double_blocks_to_swap: - mm.soft_empty_cache() + #mm.soft_empty_cache() block.to(self.offload_device, non_blocking=True) # Merge txt and img to pass through single stream blocks. @@ -692,7 +693,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin): x = block(*single_block_args) if b >= 0 and b <= self.single_blocks_to_swap: - mm.soft_empty_cache() + #mm.soft_empty_cache() block.to(self.offload_device, non_blocking=True) img = x[:, :img_seq_len, ...] diff --git a/nodes.py b/nodes.py index c5aa6dc..a82f552 100644 --- a/nodes.py +++ b/nodes.py @@ -157,6 +157,7 @@ class HyVideoModelLoader: in_channels=in_channels, out_channels=out_channels, attention_mode=attention_mode, + main_device=device, offload_device=offload_device, **HUNYUAN_VIDEO_CONFIG, **factor_kwargs @@ -646,8 +647,15 @@ class HyVideoSampler: # ) if any(q in model["quantization"] for q in ("e4m3fn", "GGUF")) else nullcontext() #with autocast_context: if model["block_swap_args"] is not None: - model["pipe"].transformer.to(device) + for name, param in model["pipe"].transformer.named_parameters(): + #print(name, param.data.device) + if "single" not in name and "double" not in name: + param.data = param.data.to(device) + model["pipe"].transformer.block_swap(model["block_swap_args"]["double_blocks_to_swap"] , model["block_swap_args"]["single_blocks_to_swap"]) + # for name, param in model["pipe"].transformer.named_parameters(): + # print(name, param.data.device) + elif model["manual_offloading"]: model["pipe"].transformer.to(device)