diff --git a/adv_control/control_sparsectrl.py b/adv_control/control_sparsectrl.py index 1267942..145abc7 100644 --- a/adv_control/control_sparsectrl.py +++ b/adv_control/control_sparsectrl.py @@ -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) diff --git a/pyproject.toml b/pyproject.toml index f8eba83..86ccc57 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 = []