add strength and steps specification

This commit is contained in:
Kohya S
2023-08-21 20:18:12 +09:00
parent 5ed3ea9c11
commit dabff295c6
3 changed files with 61 additions and 10 deletions
+6
View File
@@ -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` でリサイズしてください。
+9 -5
View File
@@ -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
]
}
],
+46 -5
View File
@@ -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)