Load textencoder to RAM by default

This commit is contained in:
kijai
2025-02-02 16:49:26 +02:00
parent e34492ad47
commit d74561122c
+12 -5
View File
@@ -634,6 +634,7 @@ class DownloadAndLoadHyVideoTextEncoder:
"apply_final_norm": ("BOOLEAN", {"default": False}),
"hidden_state_skip_layer": ("INT", {"default": 2}),
"quantization": (['disabled', 'bnb_nf4', "fp8_e4m3fn"], {"default": 'disabled'}),
"load_device": (["main_device", "offload_device"], {"default": "offload_device"}),
}
}
@@ -643,7 +644,7 @@ class DownloadAndLoadHyVideoTextEncoder:
CATEGORY = "HunyuanVideoWrapper"
DESCRIPTION = "Loads Hunyuan text_encoder model from 'ComfyUI/models/LLM'"
def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, quantization="disabled"):
def loadmodel(self, llm_model, clip_model, precision, apply_final_norm=False, hidden_state_skip_layer=2, quantization="disabled", load_device="offload_device"):
lm_type_mapping = {
"Kijai/llava-llama-3-8b-text-encoder-tokenizer": "llm",
"xtuner/llava-llama-3-8b-v1_1-transformers": "vlm",
@@ -651,6 +652,9 @@ class DownloadAndLoadHyVideoTextEncoder:
lm_type = lm_type_mapping[llm_model]
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
text_encoder_load_device = device if load_device == "main_device" else offload_device
dtype = {"bf16": torch.bfloat16, "fp16": torch.float16, "fp32": torch.float32}[precision]
quantization_config = None
if quantization == "bnb_nf4":
@@ -682,7 +686,7 @@ class DownloadAndLoadHyVideoTextEncoder:
text_encoder_precision=precision,
tokenizer_type="clipL",
logger=log,
device=device,
device=text_encoder_load_device,
)
else:
text_encoder_2 = None
@@ -706,7 +710,7 @@ class DownloadAndLoadHyVideoTextEncoder:
hidden_state_skip_layer=hidden_state_skip_layer,
apply_final_norm=apply_final_norm,
logger=log,
device=device,
device=text_encoder_load_device,
dtype=dtype,
quantization_config=quantization_config
)
@@ -1401,7 +1405,8 @@ class HyVideoEncode:
"auto_tile_size": ("BOOLEAN", {"default": True, "tooltip": "Automatically set tile size based on defaults, above settings are ignored"}),
},
"optional": {
"noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation"}),
"noise_aug_strength": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.001, "tooltip": "Strength of noise augmentation, helpful for leapfusion 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 leapfusion I2V where lower values allow for more motion"}),
}
}
@@ -1410,7 +1415,7 @@ class HyVideoEncode:
FUNCTION = "encode"
CATEGORY = "HunyuanVideoWrapper"
def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size, noise_aug_strength=0.0):
def encode(self, vae, image, enable_vae_tiling, temporal_tiling_sample_size, auto_tile_size, spatial_tile_sample_min_size, noise_aug_strength=0.0, latent_strength=1.0):
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
@@ -1437,6 +1442,8 @@ class HyVideoEncode:
if enable_vae_tiling:
vae.enable_tiling()
latents = vae.encode(image).latent_dist.sample(generator)
if latent_strength != 1.0:
latents *= latent_strength
#latents = latents * vae.config.scaling_factor
vae.to(offload_device)
print("encoded latents shape",latents.shape)