fixup block_swap start

This commit is contained in:
kijai
2024-12-04 03:53:51 +02:00
parent 3c64bbe0b8
commit 3197107b5e
2 changed files with 15 additions and 6 deletions
+6 -5
View File
@@ -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, ...]
+9 -1
View File
@@ -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)