fixup block_swap start
This commit is contained in:
@@ -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, ...]
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user