Compare commits

...
Author SHA1 Message Date
Dr.Lt.Data c48161bc0d wip 2024-03-03 12:05:58 +09:00
Dr.Lt.Data 860fead58e wip 2024-03-03 02:05:27 +09:00
Dr.Lt.Data 64c834609f wip 2024-03-03 00:10:02 +09:00
Dr.Lt.Data 8f1a4accf1 wip 2024-02-28 06:25:07 +09:00
Dr.Lt.Data 2072e8b4f1 wip 2024-02-27 12:46:28 +09:00
Dr.Lt.Data df8d82e2a7 wip 2024-02-27 11:41:51 +09:00
Dr.Lt.Data 11c65a04f9 wip: StableCascade_DetailerHookProvider 2024-02-27 11:38:30 +09:00
6 changed files with 121 additions and 6 deletions
+1
View File
@@ -182,6 +182,7 @@ NODE_CLASS_MAPPINGS = {
"UnsamplerHookProvider": UnsamplerHookProvider, "UnsamplerHookProvider": UnsamplerHookProvider,
"CoreMLDetailerHookProvider": CoreMLDetailerHookProvider, "CoreMLDetailerHookProvider": CoreMLDetailerHookProvider,
"PreviewDetailerHookProvider": PreviewDetailerHookProvider, "PreviewDetailerHookProvider": PreviewDetailerHookProvider,
"StableCascade_DetailerHookProvider": StableCascade_DetailerHookProvider,
"DetailerHookCombine": DetailerHookCombine, "DetailerHookCombine": DetailerHookCombine,
"NoiseInjectionDetailerHookProvider": NoiseInjectionDetailerHookProvider, "NoiseInjectionDetailerHookProvider": NoiseInjectionDetailerHookProvider,
+1 -1
View File
@@ -2,7 +2,7 @@ import configparser
import os import os
version_code = [4, 80] version_code = [4, 81]
version = f"V{version_code[0]}.{version_code[1]}" + (f'.{version_code[2]}' if len(version_code) > 2 else '') version = f"V{version_code[0]}.{version_code[1]}" + (f'.{version_code[2]}' if len(version_code) > 2 else '')
dependency_version = 20 dependency_version = 20
+31 -3
View File
@@ -18,6 +18,7 @@ from comfy import model_management
from impact import utils from impact import utils
from impact import impact_sampling from impact import impact_sampling
from concurrent.futures import ThreadPoolExecutor from concurrent.futures import ThreadPoolExecutor
from comfy.ldm.cascade.stage_c_coder import StageC_coder
SEG = namedtuple("SEG", SEG = namedtuple("SEG",
@@ -214,6 +215,20 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max
new_w = w new_w = w
new_h = h new_h = h
is_stable_cascade_mode = isinstance(vae.first_stage_model, StageC_coder)
if is_stable_cascade_mode:
dw = new_w % 8
dh = new_h % 8
# preserve aspect ratio as possible
if dw > 3 or dh > 3:
new_w += 8 - dw
new_h += 8 - dh
elif dw > 0 or dh > 0:
new_w -= dw
new_h -= dh
if detailer_hook is not None: if detailer_hook is not None:
new_w, new_h = detailer_hook.touch_scaled_size(new_w, new_h) new_w, new_h = detailer_hook.touch_scaled_size(new_w, new_h)
@@ -232,7 +247,14 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max
if noise_mask is not None and inpaint_model: if noise_mask is not None and inpaint_model:
positive, negative, latent_image = nodes.InpaintModelConditioning().encode(positive, negative, upscaled_image, vae, noise_mask) positive, negative, latent_image = nodes.InpaintModelConditioning().encode(positive, negative, upscaled_image, vae, noise_mask)
else: else:
latent_image = to_latent_image(upscaled_image, vae) if is_stable_cascade_mode:
latent_image = detailer_hook.stable_cascade_vae_encode(vae, upscaled_image)
if latent_image is None:
print(f"[Impact Pack] When using the StableCascade model, it is necessary to connect the StableCascade_DetailerHook.")
raise Exception("StableCascade_DetailerHook is not provided.")
else:
latent_image = to_latent_image(upscaled_image, vae)
if noise_mask is not None: if noise_mask is not None:
latent_image['noise_mask'] = noise_mask latent_image['noise_mask'] = noise_mask
@@ -258,11 +280,17 @@ def enhance_detail(image, model, clip, vae, guide_size, guide_size_for_bbox, max
refined_latent = impact_sampling.ksampler_wrapper(model2, seed2, steps2, cfg2, sampler_name2, scheduler2, positive2, negative2, refined_latent = impact_sampling.ksampler_wrapper(model2, seed2, steps2, cfg2, sampler_name2, scheduler2, positive2, negative2,
refined_latent, denoise2, refiner_ratio, refiner_model, refiner_clip, refiner_positive, refiner_negative) refined_latent, denoise2, refiner_ratio, refiner_model, refiner_clip, refiner_positive, refiner_negative)
# non-latent downscale - latent downscale cause bad quality
if detailer_hook is not None: if detailer_hook is not None:
refined_latent = detailer_hook.pre_decode(refined_latent) refined_latent = detailer_hook.pre_decode(refined_latent)
stage_b = detailer_hook.stable_cascade_stage_b(image, positive, negative, refined_latent)
else:
stage_b = None
# non-latent downscale - latent downscale cause bad quality if stage_b is None:
refined_image = vae.decode(refined_latent['samples']) refined_image = vae.decode(refined_latent['samples'])
else:
refined_image = stage_b
if detailer_hook is not None: if detailer_hook is not None:
refined_image = detailer_hook.post_decode(refined_image) refined_image = detailer_hook.post_decode(refined_image)
+28
View File
@@ -1,6 +1,7 @@
import sys import sys
from . import hooks from . import hooks
from . import defs from . import defs
import comfy
class SEGSOrderedFilterDetailerHookProvider: class SEGSOrderedFilterDetailerHookProvider:
@@ -81,3 +82,30 @@ class PreviewDetailerHookProvider:
def doit(self, quality, unique_id): def doit(self, quality, unique_id):
hook = hooks.PreviewDetailerHook(unique_id, quality) hook = hooks.PreviewDetailerHook(unique_id, quality)
return (hook, ) return (hook, )
class StableCascade_DetailerHookProvider:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"b_model": ("MODEL",),
"b_vae": ("VAE",),
"b_seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"b_steps": ("INT", {"default": 5, "min": 1, "max": 10000}),
"b_cfg": ("FLOAT", {"default": 1.1, "min": 0.0, "max": 100.0, "step": 0.1, "round": 0.01}),
"b_sampler_name": (comfy.samplers.KSampler.SAMPLERS,),
"b_scheduler": (comfy.samplers.KSampler.SCHEDULERS,),
"c_compression": ("INT", {"default": 42, "min": 4, "max": 128, "step": 1}),
},
}
RETURN_TYPES = ("DETAILER_HOOK", )
FUNCTION = "doit"
CATEGORY = "ImpactPack/Util"
def doit(self, b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, c_compression):
hook = hooks.StableCascade_DetailerHook(b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, c_compression)
return (hook, )
+56
View File
@@ -1,4 +1,6 @@
import copy import copy
import comfy_extras.nodes_stable_cascade
import nodes import nodes
from impact import utils from impact import utils
@@ -8,6 +10,8 @@ from server import PromptServer
import asyncio import asyncio
import folder_paths import folder_paths
import os import os
from impact import impact_sampling
class PixelKSampleHook: class PixelKSampleHook:
cur_step = 0 cur_step = 0
@@ -101,6 +105,20 @@ class DetailerHookCombine(PixelKSampleHookCombine):
image = self.hook2.post_paste(image) image = self.hook2.post_paste(image)
return image return image
def stable_cascade_vae_encode(self, vae, pixels):
latent = self.hook1.stable_cascade_vae_encode(vae, pixels)
if latent is not None:
return latent
return self.hook2.stable_cascade_vae_encode(vae, pixels)
def stable_cascade_stage_b(self, image, positive, negative, latent):
image = self.hook1.stable_cascade_stage_b(image, positive, negative, latent)
if image is not None:
return image
return self.hook2.stable_cascade_stage_b(image, positive, negative, latent)
class SimpleCfgScheduleHook(PixelKSampleHook): class SimpleCfgScheduleHook(PixelKSampleHook):
target_cfg = 0 target_cfg = 0
@@ -162,6 +180,44 @@ class DetailerHook(PixelKSampleHook):
def post_paste(self, image): def post_paste(self, image):
return image return image
def stable_cascade_vae_encode(self, vae, pixels):
return None
def stable_cascade_stage_b(self, image, positive, negative, latent):
return None
class StableCascade_DetailerHook(DetailerHook):
def __init__(self, b_model, b_vae, b_seed, b_steps, b_cfg, b_sampler_name, b_scheduler, c_compression):
super().__init__()
self.b_model = b_model
self.b_vae = b_vae
self.b_seed = b_seed
self.b_steps = b_steps
self.b_cfg = b_cfg
self.b_sampler_name = b_sampler_name
self.b_scheduler = b_scheduler
self.c_compression = c_compression
self.b_latent = None
def stable_cascade_vae_encode(self, vae, pixels):
obj = comfy_extras.nodes_stable_cascade.StableCascade_StageC_VAEEncode()
stage_c, stage_b = obj.generate(pixels, vae, compression=self.c_compression)
self.b_latent = stage_b
return stage_c
def stable_cascade_stage_b(self, image, positive, negative, latent):
# prepare stage_b
# self.b_latent['noise_mask'] = latent['noise_mask']
b_positive = comfy_extras.nodes_stable_cascade.StableCascade_StageB_Conditioning().set_prior(positive, latent)[0]
# stage_b sampling
b_latent = impact_sampling.ksampler_wrapper(self.b_model, self.b_seed, self.b_steps, self.b_cfg, self.b_sampler_name, self.b_scheduler, b_positive, negative, self.b_latent, 1.0)
# stage_b decoding
self.b_latent = None
return self.b_vae.decode(b_latent['samples'])
class SimpleDetailerDenoiseSchedulerHook(DetailerHook): class SimpleDetailerDenoiseSchedulerHook(DetailerHook):
def __init__(self, target_denoise): def __init__(self, target_denoise):
+4 -2
View File
@@ -8,6 +8,8 @@ from . import config
from PIL import Image, ImageFilter from PIL import Image, ImageFilter
from scipy.ndimage import zoom from scipy.ndimage import zoom
import comfy import comfy
import comfy.ldm.cascade as cascade
from comfy_extras import nodes_stable_cascade
class TensorBatchBuilder: class TensorBatchBuilder:
@@ -490,7 +492,7 @@ def crop_image(image, crop_region):
return crop_tensor4(image, crop_region) return crop_tensor4(image, crop_region)
def to_latent_image(pixels, vae): def to_latent_image(pixels, vae, compression=None):
x = pixels.shape[1] x = pixels.shape[1]
y = pixels.shape[2] y = pixels.shape[2]
if pixels.shape[1] != x or pixels.shape[2] != y: if pixels.shape[1] != x or pixels.shape[2] != y:
@@ -500,7 +502,7 @@ def to_latent_image(pixels, vae):
if hasattr(nodes.VAEEncode, "vae_encode_crop_pixels"): if hasattr(nodes.VAEEncode, "vae_encode_crop_pixels"):
# backward compatibility # backward compatibility
print(f"[Impact Pack] ComfyUI is outdated.") print(f"[Impact Pack] ComfyUI is outdated.")
pixels = nodes.VAEEncode.vae_encode_crop_pixels(pixels) pixels = vae_encode.vae_encode_crop_pixels(pixels)
t = vae.encode(pixels[:, :, :, :3]) t = vae.encode(pixels[:, :, :, :3])
return {"samples": t} return {"samples": t}