Reset Tiled Diffusion on first sample (#9)

This commit is contained in:
shiimizu
2024-01-30 14:07:33 -08:00
parent 195fed535d
commit 529a9ae578
2 changed files with 42 additions and 3 deletions
+13 -2
View File
@@ -138,6 +138,17 @@ class AbstractDiffusion:
self.weights = None
self.imagescale = ImageScale()
def reset(self):
tile_width = self.tile_width
tile_height = self.tile_height
tile_overlap = self.tile_overlap
tile_batch_size = self.tile_batch_size
self.__init__()
self.tile_width = tile_width
self.tile_height = tile_height
self.tile_overlap = tile_overlap
self.tile_batch_size = tile_batch_size
def repeat_tensor(self, x:Tensor, n:int, concat=False, concat_to=0) -> Tensor:
''' repeat the tensor on it's first dim '''
if n == 1: return x
@@ -292,11 +303,11 @@ class AbstractDiffusion:
if len(self.batched_bboxes) >= len(self.control_tensor_batch[param_id]):
self.control_tensor_batch[param_id].extend([[] for _ in range(len(self.batched_bboxes))])
# if statement: eager eval, first time when cond_hint is None.
# if statement: eager eval. first time when cond_hint is None.
if self.refresh or control.cond_hint is None or not isinstance(self.control_tensor_batch[param_id][batch_id], Tensor):
if isinstance(control, ControlNet):
dtype = control.manual_cast_dtype if control.manual_cast_dtype is not None else control.control_model.dtype
control.cond_hint = comfy.utils.common_upscale(control.cond_hint_original, PW, PH, 'nearest-exact', 'center').to(dtype).to(control.device)
control.cond_hint = comfy.utils.common_upscale(control.cond_hint_original, PW, PH, 'nearest-exact', 'center').to(dtype=dtype, device=control.device)
elif isinstance(control, T2IAdapter):
width, height = control.scale_image_to(PW, PH)
control.cond_hint = comfy.utils.common_upscale(control.cond_hint_original, width, height, 'nearest-exact', "center").float().to(control.device)
+29 -1
View File
@@ -66,7 +66,35 @@ def hook_samplers_pre_run_control():
"dedent": False,
"target_line": "if 'control' in x:",
"code_to_insert": """ try: x['control'].cleanup()\n except: ..."""
}]
},
{
"target_line": "s = model.model_sampling",
"code_to_insert": """
def find_outer_instance(target:str, target_type):
import inspect
frame = inspect.currentframe()
i = 0
while frame and i < 7:
if (found:=frame.f_locals.get(target, None)) is not None:
if isinstance(found, target_type):
return found
frame = frame.f_back
i += 1
return None
from comfy.model_patcher import ModelPatcher
if (model:=find_outer_instance('model', ModelPatcher)) is not None:
if (model_function_wrapper:=model.model_options.get('model_function_wrapper', None)) is not None:
import sys
tiled_diffusion = sys.modules.get('ComfyUI-TiledDiffusion.tiled_diffusion', None)
if tiled_diffusion is None:
for key in sys.modules:
if 'tiled_diffusion' in key:
tiled_diffusion = sys.modules[key]
break
if (AbstractDiffusion:=getattr(tiled_diffusion, 'AbstractDiffusion', None)) is not None:
if isinstance(model_function_wrapper, AbstractDiffusion):
model_function_wrapper.reset()
"""}]
fn = inject_code(pre_run_control, payload, 'a')
return create_hook(fn, 'comfy.samplers')