refactor for advanced/simple sampler, readme

This commit is contained in:
BlenderNeko
2023-05-13 02:04:12 +02:00
parent 44b937c353
commit d1353b695e
7 changed files with 246 additions and 167 deletions
+51 -8
View File
@@ -1,23 +1,66 @@
# WIP tiled sampling for ComfyUI
# Tiled sampling for ComfyUI
![panorama of the ocean, sailboats and large moody clouds](https://github.com/BlenderNeko/ComfyUI_TiledKSampler/blob/master/examples/ComfyUI_02010_.png)
this repo contains a tiled sampler for [ComfyUI](https://github.com/comfyanonymous/ComfyUI). It allows for denoising larger images by splitting it up into smaller tiles and denoising these. It tries to minimize any seams for showing up in the end result by gradually denoising all tiles one step at the time and randomizing tile positions for every step.
### settings
The tiled sampler comes with some additional settings to further control it's behavior:
The tiled samplers comes with some additional settings to further control it's behavior:
- **tile_width**: the width of the tiles.
- **tile_height**: the height of the tiles.
- **concurrent_tiles**: determines how many tiles to try and denoise concurrently.
- **tiling_strategy**: how to do the tiling
If results look tiled, it might help to increase the number of steps and to use an ancestral sampler
## Tiling strategies
roadmap:
### random:
The random tiling strategy aims to reduce the presence of seams as much as possible by slowly denoising the entire image step by step, randomizing the tile positions for each step. It does this by alternating between horizontal and vertical brick patterns, randomly offsetting the pattern each time. As the number of steps grows to infinity the strength of seams shrinks to zero. Although this random offset eliminates seams, it comes at the cost of additional overhead per step and makes this strategy incompatible with uni samplers.
<details>
<summary>
visual explanation
</summary>
![gif showing of the random brick tiling](https://github.com/BlenderNeko/ComfyUI_TiledKSampler/blob/master/examples/tiled_random.gif)
</details>
<details>
<summary>
example seamless image
</summary>
This tiling strategy is exceptionally good in hiding seams, even when starting off from complete noise, repetitions are visible but seams are not.
![gif showing of the random brick tiling](https://github.com/BlenderNeko/ComfyUI_TiledKSampler/blob/master/examples/ComfyUI_02006_.png)
</details>
### padded:
The padded tiling strategy tries to reduce seams by giving each tile more context of its surroundings through padding. It does this by further dividing each tile into 9 smaller tiles, which are denoised in such a way that a tile is always surrounded by static contex during denoising. This strategy is more prone to seams but because the location of the tiles is static, this strategy is compatible with uni samplers and has no overhead between steps. However the padding makes it so that up to 4 times as many tiles have to be denoised.
<details>
<summary>
visual explanation
</summary>
![gif showing of padded tiling](https://github.com/BlenderNeko/ComfyUI_TiledKSampler/blob/master/examples/tiled_padding.gif)
</details>
### simple
The simple tiling strategy divides the image into a static grid of tiles and denoises these one by one.
### roadmap:
- [x] latent masks
- [x] image wide control nets
- [x] T2I adaptors (requires forked version of comfy for now)
- [ ] tile wide control nets and T2I adaptors (e.g. style models)
- [ ] area conditioning
- [ ] area mask conditioning
- [ ] GLIGEN
- [x] area conditioning
- [x] area mask conditioning
- [x] GLIGEN
note:
max supported batch size is currently 1, node is currently not fully compatible with tome. These things should change after a PR to comfy
+2 -2
View File
@@ -1,3 +1,3 @@
from .nodes import NODE_CLASS_MAPPINGS
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
__all__ = ['NODE_CLASS_MAPPINGS']
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
Binary file not shown.

After

Width:  |  Height:  |  Size: 7.1 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 4.1 MiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 42 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 32 KiB

+193 -157
View File
@@ -79,6 +79,192 @@ def slices_T2I(h, h_len, w, w_len, model:comfy.sd.T2IAdapter, img):
img = model.cond_hint_original
model.cond_hint = tiling.get_slice(img, h*8, h_len*8, w*8, w_len*8).float().to(model.device)
def sample_common(model, add_noise, noise_seed, tile_width, tile_height, tiling_strategy, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, denoise=1.0):
end_at_step = min(end_at_step, steps)
device = comfy.model_management.get_torch_device()
samples = latent_image["samples"]
noise_mask = latent_image["noise_mask"] if "noise_mask" in latent_image else None
force_full_denoise = return_with_leftover_noise == "enable"
if add_noise == "disable":
noise = torch.zeros(samples.size(), dtype=samples.dtype, layout=samples.layout, device="cpu")
else:
skip = latent_image["batch_index"] if "batch_index" in latent_image else 0
noise = comfy.sample.prepare_noise(samples, noise_seed, skip)
if noise_mask is not None:
noise_mask = comfy.sample.prepare_mask(noise_mask, noise.shape, device)
shape = samples.shape
real_model = None
comfy.model_management.load_model_gpu(model)
real_model = model.model
samples = samples.to(device)
models = comfy.sample.load_additional_models(positive, negative)
sampler = comfy.samplers.KSampler(real_model, steps=steps, device=device, sampler=sampler_name, scheduler=scheduler, denoise=denoise, model_options=model.model_options)
if tiling_strategy != 'padded':
if noise_mask is not None:
samples += sampler.sigmas[start_at_step] * noise_mask * noise.to(device)
else:
samples += sampler.sigmas[start_at_step] * noise.to(device)
#cnets
cnets = [m for m in models if isinstance(m, comfy.sd.ControlNet)]
cnet_imgs = [
torch.nn.functional.interpolate(m.cond_hint_original, (shape[-2] * 8, shape[-1] * 8), mode='nearest-exact').to('cpu')
if m.cond_hint_original.shape[-2] != shape[-2] * 8 or m.cond_hint_original.shape[-1] != shape[-1] * 8 else None
for m in cnets]
#T2I
T2Is = [m for m in models if isinstance(m, comfy.sd.T2IAdapter)]
T2I_imgs = [
torch.nn.functional.interpolate(m.cond_hint_original, (shape[-2] * 8, shape[-1] * 8), mode='nearest-exact').to('cpu')
if m.cond_hint_original.shape[-2] != shape[-2] * 8 or m.cond_hint_original.shape[-1] != shape[-1] * 8 or (m.channels_in == 1 and m.cond_hint_original.shape[1] != 1) else None
for m in T2Is
]
T2I_imgs = [
torch.mean(img, 1, keepdim=True) if img is not None and m.channels_in == 1 and m.cond_hint_original.shape[1] else img
for m, img in zip(T2Is, T2I_imgs)
]
#cond area and mask
spatial_conds_pos = [
(c[1]['area'] if 'area' in c[1] else None,
comfy.sample.prepare_mask(c[1]['mask'], shape, device) if 'mask' in c[1] else None)
for c in positive
]
spatial_conds_neg = [
(c[1]['area'] if 'area' in c[1] else None,
comfy.sample.prepare_mask(c[1]['mask'], shape, device) if 'mask' in c[1] else None)
for c in negative
]
#gligen
gligen_pos = [
c[1]['gligen'] if 'gligen' in c[1] else None
for c in positive
]
gligen_neg = [
c[1]['gligen'] if 'gligen' in c[1] else None
for c in negative
]
positive_copy = comfy.sample.broadcast_cond(positive, shape[0], device)
negative_copy = comfy.sample.broadcast_cond(negative, shape[0], device)
gen = torch.manual_seed(noise_seed)
if tiling_strategy == 'random':
tiles = tiling.get_tiles_and_masks_rgrid(end_at_step - start_at_step, samples.shape, tile_height, tile_width, gen)
elif tiling_strategy == 'padded':
tiles = tiling.get_tiles_and_masks_padded(end_at_step - start_at_step, samples.shape, tile_height, tile_width)
else:
tiles = tiling.get_tiles_and_masks_simple(end_at_step - start_at_step, samples.shape, tile_height, tile_width)
total_steps = sum([num_steps for img_pass in tiles for steps_list in img_pass for _,_,_,_,num_steps,_ in steps_list])
current_step = [0]
with tqdm(total=total_steps) as pbar_tqdm:
pbar = comfy.utils.ProgressBar(total_steps)
def callback(step, x0, x, total_steps):
current_step[0] += 1
pbar.update_absolute(current_step[0])
pbar_tqdm.update(1)
for img_pass in tiles:
for i in range(len(img_pass)):
for tile_h, tile_h_len, tile_w, tile_w_len, tile_steps, tile_mask in img_pass[i]:
#if we have masks get mask slices and see if we can skip slices
if noise_mask is not None or tile_mask is not None:
if noise_mask is not None:
tiled_mask = tiling.get_slice(noise_mask, tile_h, tile_h_len, tile_w, tile_w_len)
if tile_mask is not None:
tiled_mask *= tile_mask.to(device)
else:
tiled_mask = tile_mask.to(device)
if tiled_mask.sum().cpu() == 0.0:
continue
else:
tiled_mask = None
tiled_latent = tiling.get_slice(samples, tile_h, tile_h_len, tile_w, tile_w_len)
if tiling_strategy == 'padded':
tiled_noise = tiling.get_slice(noise, tile_h, tile_h_len, tile_w, tile_w_len).to(device)
else:
if tiled_mask is None:
tiled_noise = torch.zeros_like(tiled_latent)
else:
tiling.get_slice(noise, tile_h, tile_h_len, tile_w, tile_w_len).to(device) * (1 - tiled_mask)
#TODO: all other condition based stuff like area sets and GLIGEN should also happen here
#cnets
for m, img in zip(cnets, cnet_imgs):
slice_cnet(tile_h, tile_h_len, tile_w, tile_w_len, m, img)
#T2I
for m, img in zip(T2Is, T2I_imgs):
slices_T2I(tile_h, tile_h_len, tile_w, tile_w_len, m, img)
pos = copy_cond(positive_copy)
neg = copy_cond(negative_copy)
#cond areas
pos = [slice_cond(tile_h, tile_h_len, tile_w, tile_w_len, c, area) for c, area in zip(pos, spatial_conds_pos)]
pos = [c for c, ignore in pos if not ignore]
neg = [slice_cond(tile_h, tile_h_len, tile_w, tile_w_len, c, area) for c, area in zip(neg, spatial_conds_neg)]
neg = [c for c, ignore in neg if not ignore]
#gligen
for (_, cond), gligen in zip(pos, gligen_pos):
slice_gligen(tile_h, tile_h_len, tile_w, tile_w_len, cond, gligen)
for (_, cond), gligen in zip(neg, gligen_neg):
slice_gligen(tile_h, tile_h_len, tile_w, tile_w_len, cond, gligen)
tile_result = sampler.sample(tiled_noise, pos, neg, cfg=cfg, latent_image=tiled_latent, start_step=start_at_step + i * tile_steps, last_step=start_at_step + i*tile_steps + tile_steps, force_full_denoise=force_full_denoise and i+1 == end_at_step - start_at_step, denoise_mask=tiled_mask, callback=callback, disable_pbar=True)
tiling.set_slice(samples, tile_result, tile_h, tile_h_len, tile_w, tile_w_len, tiled_mask)
comfy.sample.cleanup_additional_models(models)
out = latent_image.copy()
out["samples"] = samples.cpu()
return (out, )
class TiledKSampler:
@classmethod
def INPUT_TYPES(s):
return {"required":
{"model": ("MODEL",),
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
"tile_width": ("INT", {"default": 512, "min": 256, "max": MAX_RESOLUTION, "step": 64}),
"tile_height": ("INT", {"default": 512, "min": 256, "max": MAX_RESOLUTION, "step": 64}),
"tiling_strategy": (["random", "padded", 'simple'], ),
"steps": ("INT", {"default": 20, "min": 1, "max": 10000}),
"cfg": ("FLOAT", {"default": 8.0, "min": 0.0, "max": 100.0}),
"sampler_name": (comfy.samplers.KSampler.SAMPLERS, ),
"scheduler": (comfy.samplers.KSampler.SCHEDULERS, ),
"positive": ("CONDITIONING", ),
"negative": ("CONDITIONING", ),
"latent_image": ("LATENT", ),
"denoise": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}),
}}
RETURN_TYPES = ("LATENT",)
FUNCTION = "sample"
CATEGORY = "sampling"
def sample(self, model, seed, tile_width, tile_height, tiling_strategy, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise):
steps_total = int(steps / denoise)
return sample_common(model, 'enable', seed, tile_width, tile_height, tiling_strategy, steps_total, cfg, sampler_name, scheduler, positive, negative, latent_image, steps_total-steps, steps_total, 'disable', denoise=1.0)
class TiledKSamplerAdvanced:
@classmethod
def INPUT_TYPES(s):
@@ -106,167 +292,17 @@ class TiledKSamplerAdvanced:
CATEGORY = "sampling"
def sample(self, model, add_noise, noise_seed, tile_width, tile_height, tiling_strategy, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, denoise=1.0):
end_at_step = min(end_at_step, steps)
device = comfy.model_management.get_torch_device()
samples = latent_image["samples"]
noise_mask = latent_image["noise_mask"] if "noise_mask" in latent_image else None
force_full_denoise = return_with_leftover_noise == "enable"
if add_noise == "disable":
noise = torch.zeros(samples.size(), dtype=samples.dtype, layout=samples.layout, device="cpu")
else:
skip = latent_image["batch_index"] if "batch_index" in latent_image else 0
noise = comfy.sample.prepare_noise(samples, noise_seed, skip)
if noise_mask is not None:
noise_mask = comfy.sample.prepare_mask(noise_mask, noise.shape, device)
shape = samples.shape
real_model = None
comfy.model_management.load_model_gpu(model)
real_model = model.model
samples = samples.to(device)
models = comfy.sample.load_additional_models(positive, negative)
sampler = comfy.samplers.KSampler(real_model, steps=steps, device=device, sampler=sampler_name, scheduler=scheduler, denoise=denoise, model_options=model.model_options)
if tiling_strategy != 'padded':
if noise_mask is not None:
samples += sampler.sigmas[start_at_step] * noise_mask * noise.to(device)
else:
samples += sampler.sigmas[start_at_step] * noise.to(device)
#cnets
cnets = [m for m in models if isinstance(m, comfy.sd.ControlNet)]
cnet_imgs = [
torch.nn.functional.interpolate(m.cond_hint_original, (shape[-2] * 8, shape[-1] * 8), mode='nearest-exact').to('cpu')
if m.cond_hint_original.shape[-2] != shape[-2] * 8 or m.cond_hint_original.shape[-1] != shape[-1] * 8 else None
for m in cnets]
#T2I
T2Is = [m for m in models if isinstance(m, comfy.sd.T2IAdapter)]
T2I_imgs = [
torch.nn.functional.interpolate(m.cond_hint_original, (shape[-2] * 8, shape[-1] * 8), mode='nearest-exact').to('cpu')
if m.cond_hint_original.shape[-2] != shape[-2] * 8 or m.cond_hint_original.shape[-1] != shape[-1] * 8 or (m.channels_in == 1 and m.cond_hint_original.shape[1] != 1) else None
for m in T2Is
]
T2I_imgs = [
torch.mean(img, 1, keepdim=True) if img is not None and m.channels_in == 1 and m.cond_hint_original.shape[1] else img
for m, img in zip(T2Is, T2I_imgs)
]
#cond area and mask
spatial_conds_pos = [
(c[1]['area'] if 'area' in c[1] else None,
comfy.sample.prepare_mask(c[1]['mask'], shape, device) if 'mask' in c[1] else None)
for c in positive
]
spatial_conds_neg = [
(c[1]['area'] if 'area' in c[1] else None,
comfy.sample.prepare_mask(c[1]['mask'], shape, device) if 'mask' in c[1] else None)
for c in negative
]
#gligen
gligen_pos = [
c[1]['gligen'] if 'gligen' in c[1] else None
for c in positive
]
gligen_neg = [
c[1]['gligen'] if 'gligen' in c[1] else None
for c in negative
]
positive_copy = comfy.sample.broadcast_cond(positive, shape[0], device)
negative_copy = comfy.sample.broadcast_cond(negative, shape[0], device)
gen = torch.manual_seed(noise_seed)
if tiling_strategy == 'random':
tiles = tiling.get_tiles_and_masks_rgrid(end_at_step - start_at_step, samples.shape, tile_height, tile_width, gen)
elif tiling_strategy == 'padded':
tiles = tiling.get_tiles_and_masks_padded(end_at_step - start_at_step, samples.shape, tile_height, tile_width)
else:
tiles = tiling.get_tiles_and_masks_simple(end_at_step - start_at_step, samples.shape, tile_height, tile_width)
total_steps = sum([num_steps for img_pass in tiles for steps_list in img_pass for _,_,_,_,num_steps,_ in steps_list])
current_step = [0]
with tqdm(total=total_steps) as pbar_tqdm:
pbar = comfy.utils.ProgressBar(total_steps)
def callback(step, x0, x, total_steps):
current_step[0] += 1
pbar.update_absolute(current_step[0])
pbar_tqdm.update(1)
for img_pass in tiles:
for i in range(len(img_pass)):
for tile_h, tile_h_len, tile_w, tile_w_len, tile_steps, tile_mask in img_pass[i]:
#if we have masks get mask slices and see if we can skip slices
if noise_mask is not None or tile_mask is not None:
if noise_mask is not None:
tiled_mask = tiling.get_slice(noise_mask, tile_h, tile_h_len, tile_w, tile_w_len)
if tile_mask is not None:
tiled_mask *= tile_mask.to(device)
else:
tiled_mask = tile_mask.to(device)
if tiled_mask.sum().cpu() == 0.0:
continue
else:
tiled_mask = None
tiled_latent = tiling.get_slice(samples, tile_h, tile_h_len, tile_w, tile_w_len)
if tiling_strategy == 'padded':
tiled_noise = tiling.get_slice(noise, tile_h, tile_h_len, tile_w, tile_w_len).to(device)
else:
if tiled_mask is None:
tiled_noise = torch.zeros_like(tiled_latent)
else:
tiling.get_slice(noise, tile_h, tile_h_len, tile_w, tile_w_len).to(device) * (1 - tiled_mask)
#TODO: all other condition based stuff like area sets and GLIGEN should also happen here
#cnets
for m, img in zip(cnets, cnet_imgs):
slice_cnet(tile_h, tile_h_len, tile_w, tile_w_len, m, img)
#T2I
for m, img in zip(T2Is, T2I_imgs):
slices_T2I(tile_h, tile_h_len, tile_w, tile_w_len, m, img)
pos = copy_cond(positive_copy)
neg = copy_cond(negative_copy)
#cond areas
pos = [slice_cond(tile_h, tile_h_len, tile_w, tile_w_len, c, area) for c, area in zip(pos, spatial_conds_pos)]
pos = [c for c, ignore in pos if not ignore]
neg = [slice_cond(tile_h, tile_h_len, tile_w, tile_w_len, c, area) for c, area in zip(neg, spatial_conds_neg)]
neg = [c for c, ignore in neg if not ignore]
#gligen
for (_, cond), gligen in zip(pos, gligen_pos):
slice_gligen(tile_h, tile_h_len, tile_w, tile_w_len, cond, gligen)
for (_, cond), gligen in zip(neg, gligen_neg):
slice_gligen(tile_h, tile_h_len, tile_w, tile_w_len, cond, gligen)
tile_result = sampler.sample(tiled_noise, pos, neg, cfg=cfg, latent_image=tiled_latent, start_step=start_at_step + i * tile_steps, last_step=start_at_step + i*tile_steps + tile_steps, force_full_denoise=force_full_denoise and i+1 == end_at_step - start_at_step, denoise_mask=tiled_mask, callback=callback, disable_pbar=True)
tiling.set_slice(samples, tile_result, tile_h, tile_h_len, tile_w, tile_w_len, tiled_mask)
comfy.sample.cleanup_additional_models(models)
out = latent_image.copy()
out["samples"] = samples.cpu()
return (out, )
return sample_common(model, add_noise, noise_seed, tile_width, tile_height, tiling_strategy, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, start_at_step, end_at_step, return_with_leftover_noise, denoise=1.0)
NODE_CLASS_MAPPINGS = {
"BNK_TiledKSamplerAdvanced": TiledKSamplerAdvanced,
"BNK_TiledKSampler": TiledKSampler,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"BNK_TiledKSamplerAdvanced": "TiledK Sampler (Advanced)",
"BNK_TiledKSampler": "Tiled KSampler",
}