add strength and steps specification
This commit is contained in:
@@ -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` でリサイズしてください。
|
||||
|
||||
@@ -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
|
||||
]
|
||||
}
|
||||
],
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user