Merge PR #170 from Kosinkadink/develop - fix fp8 support for SparseCtrl

Fix fp8 support for SparseCtrl
This commit is contained in:
Jedrzej Kosinski
2024-08-30 08:35:40 -05:00
committed by GitHub
2 changed files with 15 additions and 1 deletions
+14
View File
@@ -27,6 +27,7 @@ from comfy.ldm.modules.diffusionmodules.openaimodel import TimestepEmbedSequenti
from comfy.model_patcher import ModelPatcher
import comfy.ops
import comfy.model_management
import comfy.utils
from .logger import logger
from .utils import (BIGMAX, AbstractPreprocWrapper, disable_weight_init_clean_groupnorm,
@@ -118,6 +119,7 @@ class SparseModelPatcher(ModelPatcher):
to_return = super().load(device_to=device_to, lowvram_model_memory=lowvram_model_memory, *args, **kwargs)
if lowvram_model_memory > 0:
self._patch_lowvram_extras(device_to=device_to)
self._handle_float8_pe_tensors()
return to_return
def _patch_lowvram_extras(self, device_to=None):
@@ -138,6 +140,18 @@ class SparseModelPatcher(ModelPatcher):
if device_to is not None:
comfy.utils.set_attr(self.model.motion_wrapper, key, comfy.utils.get_attr(self.model.motion_wrapper, key).to(device_to))
def _handle_float8_pe_tensors(self):
if self.model.motion_wrapper is not None:
remaining_tensors = list(self.model.motion_wrapper.state_dict().keys())
pe_tensors = [x for x in remaining_tensors if '.pe' in x]
is_first = True
for key in pe_tensors:
if is_first:
is_first = False
if comfy.utils.get_attr(self.model.motion_wrapper, key).dtype not in [torch.float8_e5m2, torch.float8_e4m3fn]:
break
comfy.utils.set_attr(self.model.motion_wrapper, key, comfy.utils.get_attr(self.model.motion_wrapper, key).half())
# NOTE: no longer called by ComfyUI, but here for backwards compatibility
def patch_model_lowvram(self, device_to=None, *args, **kwargs):
patched_model = super().patch_model_lowvram(device_to, *args, **kwargs)
+1 -1
View File
@@ -1,7 +1,7 @@
[project]
name = "comfyui-advanced-controlnet"
description = "Nodes for scheduling ControlNet strength across timesteps and batched latents, as well as applying custom weights and attention masks."
version = "1.2.2"
version = "1.2.3"
license = { file = "LICENSE" }
dependencies = []