From dabff295c6541cf6f1ba2a44b440cc0120a51429 Mon Sep 17 00:00:00 2001 From: Kohya S Date: Mon, 21 Aug 2023 20:18:12 +0900 Subject: [PATCH] add strength and steps specification --- README.md | 6 +++++ lllite_workflow.json | 14 +++++++---- node_control_net_lllite.py | 51 ++++++++++++++++++++++++++++++++++---- 3 files changed, 61 insertions(+), 10 deletions(-) diff --git a/README.md b/README.md index 05f51ac..33053ba 100644 --- a/README.md +++ b/README.md @@ -13,6 +13,12 @@ ControlNet-LLLiteがそもそもきわめて実験的な実装のため、問題 [サンプルのワークフロー](lllite_workflow.json)を読み込んでください。 +`strength`に効果の強さを指定できます。1.0でデフォルト、0.0で効果なしです。 + +`steps`と`start_percent`、`end_percent`で拡散ステップの一部にだけ効果を適用できます。`steps`にsamplerに指定したステップ数を指定し、`start_percent`と`end_percent`にそれぞれ開始と終了のステップを0から100で指定します。 + +(ノード内で全体のステップ数を確認できないためこのような仕様になっています。具体的な適用範囲はコンソール出力を確認してください。) + # Hint + 生成画像サイズと制御用画像サイズが異なる場合はワークフローにあるように `image/upscaling/UpscaleImage` でリサイズしてください。 diff --git a/lllite_workflow.json b/lllite_workflow.json index 52d7710..eaf11ac 100644 --- a/lllite_workflow.json +++ b/lllite_workflow.json @@ -538,10 +538,10 @@ 405, 433 ], - "size": [ - 396.92310800781274, - 80.65141015625022 - ], + "size": { + "0": 396.923095703125, + "1": 174 + }, "flags": {}, "order": 8, "mode": 0, @@ -572,7 +572,11 @@ "Node name for S&R": "LLLiteLoader" }, "widgets_values": [ - "controllllite_v01032064e_sdxl_canny_anime.safetensors" + "controllllite_v01032064e_sdxl_canny_anime.safetensors", + 1, + 0, + 0, + 0 ] } ], diff --git a/node_control_net_lllite.py b/node_control_net_lllite.py index 60f888e..cc33b95 100644 --- a/node_control_net_lllite.py +++ b/node_control_net_lllite.py @@ -34,7 +34,12 @@ def extra_options_to_module_prefix(extra_options): return module_pfx -def load_control_net_lllite_patch(path, cond_image): +def load_control_net_lllite_patch(path, cond_image, multiplier, num_steps, start_percent, end_percent): + # calculate start and end step + start_step = math.floor(num_steps * start_percent * 0.01) if start_percent > 0 else 0 + end_step = math.floor(num_steps * end_percent * 0.01) if end_percent > 0 else num_steps + + # load weights ctrl_sd = comfy.utils.load_torch_file(path, safe_load=True) # split each weights for each module @@ -66,9 +71,15 @@ def load_control_net_lllite_patch(path, cond_image): depth=depth, cond_emb_dim=weights["conditioning1.0.weight"].shape[0] * 2, mlp_dim=weights["down.0.weight"].shape[0], + multiplier=multiplier, + num_steps=num_steps, + start_step=start_step, + end_step=end_step, ) info = module.load_state_dict(weights) modules[module_name] = module + if len(modules) == 1: + module.is_first = True print(f"loaded {path} successfully, {len(modules)} modules") @@ -122,10 +133,19 @@ class LLLiteModule(torch.nn.Module): depth: int, cond_emb_dim: int, mlp_dim: int, + multiplier: int, + num_steps: int, + start_step: int, + end_step: int, ): super().__init__() self.name = name self.is_conv2d = is_conv2d + self.multiplier = multiplier + self.num_steps = num_steps + self.start_step = start_step + self.end_step = end_step + self.is_first = False modules = [] modules.append(torch.nn.Conv2d(3, cond_emb_dim // 2, kernel_size=4, stride=4, padding=0)) # to latent (from VAE) size*2 @@ -172,17 +192,34 @@ class LLLiteModule(torch.nn.Module): self.depth = depth self.cond_image = None self.cond_emb = None + self.current_step = 0 # @torch.inference_mode() def set_cond_image(self, cond_image): # print("set_cond_image", self.name) self.cond_image = cond_image self.cond_emb = None + self.current_step = 0 def forward(self, x): + if self.num_steps > 0: + if self.current_step < self.start_step: + self.current_step += 1 + return torch.zeros_like(x) + elif self.current_step >= self.end_step: + if self.is_first and self.current_step == self.end_step: + print(f"end LLLite: step {self.current_step}") + self.current_step += 1 + return torch.zeros_like(x) + else: + if self.is_first and self.current_step ==self.start_step: + print(f"start LLLite: step {self.current_step}") + self.current_step += 1 + if self.current_step >= self.num_steps: + self.current_step = 0 # reset + if self.cond_emb is None: # print(f"cond_emb is None, {self.name}") - # TODO resize image here cx = self.conditioning1(self.cond_image.to(x.device, dtype=x.dtype)) if not self.is_conv2d: # reshape / b,c,h,w -> b,h*w,c @@ -204,7 +241,7 @@ class LLLiteModule(torch.nn.Module): cx = torch.cat([cx, self.down(x)], dim=1 if self.is_conv2d else 2) cx = self.mid(cx) cx = self.up(cx) - return cx + return cx * self.multiplier class LLLiteLoader: @@ -218,6 +255,10 @@ class LLLiteLoader: "model": ("MODEL",), "model_name": (get_file_list(os.path.join(CURRENT_DIR, "models")),), "cond_image": ("IMAGE",), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01}), + "steps": ("INT", {"default": 0, "min": 0, "max": 200, "step": 1}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.1}), + "end_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 100.0, "step": 0.1}), } } @@ -225,13 +266,13 @@ class LLLiteLoader: FUNCTION = "load_lllite" CATEGORY = "loaders" - def load_lllite(self, model, model_name, cond_image): + def load_lllite(self, model, model_name, cond_image, strength, steps, start_percent, end_percent): # cond_image is b,h,w,3, 0-1 model_path = os.path.join(CURRENT_DIR, os.path.join(CURRENT_DIR, "models", model_name)) model_lllite = model.clone() - patch = load_control_net_lllite_patch(model_path, cond_image) + patch = load_control_net_lllite_patch(model_path, cond_image, strength, steps, start_percent, end_percent) if patch is not None: model_lllite.set_model_attn1_patch(patch) model_lllite.set_model_attn2_patch(patch)