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 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())
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user