From 13bd60bcc9600d1a5fd6096faa95a8e0e41b26f4 Mon Sep 17 00:00:00 2001 From: shiimizu Date: Tue, 30 Jan 2024 20:35:41 -0800 Subject: [PATCH] GLIGEN fixes. --- utils.py | 18 ++++++++++-------- 1 file changed, 10 insertions(+), 8 deletions(-) diff --git a/utils.py b/utils.py index 7139a29..1dcc980 100644 --- a/utils.py +++ b/utils.py @@ -100,14 +100,16 @@ def hook_samplers_pre_run_control(): def hook_gligen__set_position(): from comfy.gligen import Gligen - payload = [{ - "target_line": "module = self.module_list[key]", - "code_to_insert": """ - nonlocal objs - if x.shape[0] > objs.shape[0]: - objs = objs.repeat(-(x.shape[0] // -objs.shape[0]),1,1) - """}] - fn = inject_code(Gligen._set_position, payload, 'a') + source=inspect.getsource(Gligen._set_position) + replace_str=""" + nonlocal objs + if x.shape[0] > objs.shape[0]: + _objs = objs.repeat(-(x.shape[0] // -objs.shape[0]),1,1) + else: + _objs = objs + return module(x, _objs)""" + modified_source = dedent(source.replace(" return module(x, objs)", replace_str, 1)) + fn = write_to_file_and_return_fn(Gligen._set_position, modified_source) return create_hook(fn, 'comfy.gligen', 'Gligen._set_position') def create_hook(fn, module_name:str, target = None, orig_key = None):