Fix FLF2V when using vram management node

This commit is contained in:
kijai
2025-04-22 09:44:46 +03:00
parent 5109a74839
commit 949c887e0c
2 changed files with 2 additions and 2 deletions
+1 -1
View File
@@ -75,7 +75,7 @@ def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict,
for name, module in model.named_children(): for name, module in model.named_children():
for source_module, target_module in module_map.items(): for source_module, target_module in module_map.items():
if isinstance(module, source_module): if isinstance(module, source_module):
if "rope_embedder" in name or "patch_embedding" in name: if "rope_embedder" in name or "patch_embedding" in name or "emb_pos" in name:
continue continue
num_param = sum(p.numel() for p in module.parameters()) num_param = sum(p.numel() for p in module.parameters())
+1 -1
View File
@@ -695,7 +695,7 @@ class MLPProj(torch.nn.Module):
def forward(self, image_embeds): def forward(self, image_embeds):
if hasattr(self, 'emb_pos'): if hasattr(self, 'emb_pos'):
image_embeds = image_embeds + self.emb_pos image_embeds = image_embeds + self.emb_pos.to(image_embeds.device)
clip_extra_context_tokens = self.proj(image_embeds) clip_extra_context_tokens = self.proj(image_embeds)
return clip_extra_context_tokens return clip_extra_context_tokens