Fix FLF2V when using vram management node
This commit is contained in:
@@ -75,7 +75,7 @@ def enable_vram_management_recursively(model: torch.nn.Module, module_map: dict,
|
||||
for name, module in model.named_children():
|
||||
for source_module, target_module in module_map.items():
|
||||
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
|
||||
|
||||
num_param = sum(p.numel() for p in module.parameters())
|
||||
|
||||
@@ -695,7 +695,7 @@ class MLPProj(torch.nn.Module):
|
||||
|
||||
def forward(self, image_embeds):
|
||||
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)
|
||||
return clip_extra_context_tokens
|
||||
|
||||
|
||||
Reference in New Issue
Block a user