4 Commits
Author SHA1 Message Date
Austin Mroz bd57dd11e3 Fix mask generation with multiple latetnts 2024-05-31 14:43:52 -05:00
Austin Mroz 1ebb9df6eb Add option for blurring mask before thresholding 2024-05-31 13:41:24 -05:00
Austin Mroz 2c0544b0c3 Add nodes to extract confidence mask
Since the method for determining areas of low confidence shows greater
promise than the method for oversampling those areas, two additional
nodes have been added, one to measure uncertainty in generation, and
another to extract a mask after generation has occurred.
2024-05-31 12:01:07 -05:00
Austin Mroz f31c016a32 Expose configuration options, add comments
Reducing the multiplier on resolution for oversampling seems to improve
output quality, but even in the best case, the effect of the node seems
little better than noise.
2024-05-30 18:18:15 -05:00
+117 -48
View File
@@ -8,21 +8,6 @@ from comfy.k_diffusion import sampling
from comfy.samplers import KSAMPLER
from tqdm.auto import trange
class LatentFFTAsImage:
"""Takes a latent as input , """
@classmethod
def INPUT_TYPES(s):
return {"required": {
"latent" : ("LATENT",),
}}
RETURN_TYPES = ("IMAGE",)
FUNCTION = "latentfft_as_image"
CATEGORY = "latent/advanced"
def latentfft_as_image(self, latent):
latent = latent['samples'].permute((0,2,3,1))
return (latent.sigmoid_(),)
def dump_image(tensor):
if (tensor == tensor.flatten()[0]).all():
print("trivial dump")
@@ -38,7 +23,7 @@ def dump_image(tensor):
Image.fromarray(b).save("test.png")
@torch.no_grad()
def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1.):
def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1., sub_weight=3, early_terminate=False, resolution_mult=1.25):
"""Implements Algorithm 2 (Euler steps) from Karras et al. (2022)."""
extra_args = {} if extra_args is None else extra_args
s_in = x.new_ones([x.shape[0]])
@@ -50,7 +35,7 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
ch = h-h//4+1
kernel = torch.ones((1,1,h//4,w//4), dtype=x.dtype, device=x.device)
prev_denoised = None
early_terminate = False
oversample_performed = False
for i in trange(len(sigmas) - 1, disable=disable):
gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0.
sigma_hat = sigmas[i] * (gamma + 1)
@@ -58,16 +43,21 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
eps = torch.randn_like(x) * s_noise
x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5
denoised = model(x, sigma_hat * s_in, **extra_args)
#Most change is happening where magnitude of dt is greatest -> blur, find peak
maskl = (i/len(sigmas)*base_mask.max()-base_mask).clip(max=1,min=0).unsqueeze(0).unsqueeze(0)
maskh = ((i+1)/len(sigmas)*base_mask.max()-base_mask).clip(max=1,min=0).unsqueeze(0).unsqueeze(0)
#A filter of where changes are expected
mask = torch.fft.ifftshift(maskh*(1-maskl))
#Just the expected changes
filt = torch.fft.ifft2(mask*torch.fft.fft2(denoised)).real
if prev_denoised is not None:
#changes made outside the current sigma schedule
ext = prev_denoised+filt-denoised
dist = (ext*ext).sum(1)/-dt
#The distance of changes relative to the current step
dist = (ext*ext).sum(1)/math.sqrt(-dt)
hist+= dist
#The sum of distances that would be modified by a subrender for each center point
summed = torch.nn.functional.conv2d(dist, kernel, groups=1)
#The most valuable of possible center points
focal = torch.unravel_index(summed.view((n,ch*cw)).argmax(1), (ch,cw))
focal = torch.stack(focal).transpose(0,1)
sublocs = []
@@ -75,34 +65,35 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
focalm = []
for j,ind in enumerate(focal):
focalm.append(summed[j, *ind])
#TODO: allow for simultaneous selection of multiple zones
#An arbitrary cutoff
if summed[j, *ind] > 1000:
subx = x[j,:,ind[0]:ind[0]+h//4,ind[1]:ind[1]+w//4]
#scale+noise subx
sublocs.append(subx)
subinds.append(j)
if len(sublocs) > 0:
early_terminate = True
oversample_performed = True
subx = torch.stack(sublocs)
#subx = subx.repeat_interleave(2, dim=2).repeat_interleave(2, dim=3)
subx = interpolate(subx, size=(h//2,w//2), mode='bicubic')
sh,sw = h//4,w//4
subx = interpolate(subx, size=(int(sh*resolution_mult),
int(sw*resolution_mult)), mode='bicubic')
#The upscaled/interpolated pixels lack higher frequency noise,
#so additional noise must be added
subx += torch.randn_like(subx)*(1-math.sqrt(2))*sigma_hat*s_in
print("doing subrender")
subdn = model(subx, sigma_hat*s_in*math.sqrt(2), **extra_args)
#downscale
subdn = interpolate(subdn, size=(h//4,w//4), mode='bicubic')
#TODO: Does added noise need to be filtered out?
subdn = interpolate(subdn, size=(sh,sw), mode='bicubic')
for j,dn in zip(subinds,subdn):
subx = denoised[j,:,focal[j][0]:focal[j][0]+h//4,focal[j][1]:focal[j][1]+w//4]
sub_weight = -1
if sub_weight == -1:
subx[:] = dn
else:
subx[:] += sub_weight*dn
subx /= sub_weight+1
focal[:,0] += h//8
focal[:,1] += w//8
print(focal*8)
print(focalm)
focal[:,0] += h//8
focal[:,1] += w//8
print(focal*8, focalm)
prev_denoised = denoised
d = sampling.to_d(x, sigma_hat, denoised)
@@ -111,40 +102,118 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
dt = sigmas[i + 1] - sigma_hat
# Euler method
x = x + d * dt
if early_terminate:
if early_terminate and oversample_performed:
break
summed = torch.nn.functional.conv2d(hist, kernel, groups=1)
dump_image(hist)
#Debugging code to dump the measured distances
#summed = torch.nn.functional.conv2d(hist, kernel, groups=1)
#dump_image(hist)
return x
class DynamicSampler:
@classmethod
def INPUT_TYPES(s):
return {"required": {}}
return {"required": {"sub_weight": ("FLOAT", {"default": 3, "min": -1, "step": 0.01}),
"resolution_mult": ("FLOAT", {"default": 1.25, "min": 1, "step": 0.01}),
"early_terminate": ("BOOLEAN", {"default": False}),}}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "get_sampler"
def get_sampler(self):
return (KSAMPLER(sample_dynamic),)
def get_sampler(self, sub_weight, early_terminate, resolution_mult):
return (KSAMPLER(lambda *args, **kwargs: sample_dynamic(*args, sub_weight=sub_weight, early_terminate=early_terminate, resolution_mult=resolution_mult, **kwargs)),)
class LatentAsImage:
"""Converts a latent to an image with minimal processing. Provides a means of viewing
latent channels visually with minimal overhead"""
class MeasuredSampler(KSAMPLER):
def __init__(self, sampler):
self.sampler = sampler.sampler_function
self.prev_denoised = None
self.prev_sigma = 0
super().__init__(self.wrapped_sample)
def wrapped_sample(self, *args, **kwargs):
original_callback = kwargs.get("callback", None)
def callback(args):
self.callback(args)
if original_callback is not None:
original_callback(args)
self.sigmas = args[2]
self.steps = len(args[2])
kwargs["callback"] = callback
x = args[1]
n,c,h,w = x.shape
self.hist = torch.zeros((n,h,w), device=x.device, dtype=torch.float32)
kx, ky = torch.meshgrid(torch.linspace(-h/2,h/2,h),torch.linspace(-w/2,w/2,w), indexing="ij")
self.base_mask = (kx*kx + ky*ky).sqrt().to(x.device)
self.prev_denoised = None
x = self.sampler(*args, **kwargs)
return x
def callback(self, args):
denoised = args["denoised"]
i = args["i"]
dt = self.sigmas[i+1] - args['sigma_hat']
if self.prev_denoised is not None:
maskl = (i/self.steps*self.base_mask.max()-self.base_mask).clip(max=1,min=0).unsqueeze(0).unsqueeze(0)
maskh = ((i+1)/self.steps*self.base_mask.max()-self.base_mask).clip(max=1,min=0).unsqueeze(0).unsqueeze(0)
#A filter of where changes are expected
mask = torch.fft.ifftshift(maskh*(1-maskl))
#Just the expected changes
filt = torch.fft.ifft2(mask*torch.fft.fft2(denoised)).real
#changes made outside the current sigma schedule
ext = self.prev_denoised+filt-denoised
#The distance of changes relative to the current step
dist = (ext*ext).sum(1)/math.sqrt(-dt)
self.hist+= dist
self.prev_denoised = denoised
#Blur functions borrowed from comfy_extras/nodes_post_processing
#with slight modifications
def gaussian_blur(latents, sigma=1):
radius = min(latents.shape[2:])//2
size = radius*2+1
maxl = size // 2
x, y = torch.meshgrid(torch.linspace(-maxl, maxl, size), torch.linspace(-maxl, maxl, size), indexing="ij")
d = (x * x + y * y) / (2 * sigma * sigma)
mask = torch.exp(-d) / (2 * torch.pi * sigma * sigma)
kernel = (mask / mask.sum())[None, None].to(latents.device)
padded_latents = torch.nn.functional.pad(latents, [radius]*4, 'reflect')
blurred = torch.nn.functional.conv2d(padded_latents, kernel, padding=(radius*2+1) // 2, groups=1)
return blurred[:, :, radius:-radius, radius:-radius]
class MeasuredSamplerNode:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"latent" : ("LATENT",),
}}
return {"required": {"sampler": ("SAMPLER",),}}
RETURN_TYPES = ("SAMPLER", "MASK_PROMISE")
CATEGORY = "sampling/custom_sampling/samplers"
RETURN_TYPES = ("IMAGE",)
FUNCTION = "latent_as_image"
CATEGORY = "latent/advanced"
def latent_as_image(self, latent):
latent = latent['samples'].permute((0,2,3,1))
return (latent.sigmoid_(),)
FUNCTION = "get_sampler"
def get_sampler(self, sampler):
s = MeasuredSampler(sampler)
return (s, s)
class ResolveMaskPromise:
@classmethod
def INPUT_TYPES(s):
return {"required": {"latent": ("LATENT",), "mask_promise": ("MASK_PROMISE",),
"upper_threshold": ("FLOAT", {"default": .8, "step": .01, "min": 0, "max": 1}),
"lower_threshold": ("FLOAT", {"default": .2, "step": .01, "min": 0, "max": 1}),
"blur_sigma": ("FLOAT", {"default": 0, "min": 0, "step": .01}),}}
RETURN_TYPES = ("MASK",)
FUNCTION = "get_mask"
def get_mask(self, latent, mask_promise, lower_threshold, upper_threshold, blur_sigma):
#NOTE: latent is only used to ensure this executes after sampling
hist = mask_promise.hist
if blur_sigma > 0:
hist = gaussian_blur(hist.unsqueeze(1), blur_sigma).squeeze(1)
sorted_hist = hist.flatten(start_dim=1).sort().values
lower = sorted_hist[:,int((sorted_hist.size(1)-1)*lower_threshold)][:,None,None]
upper = sorted_hist[:,int((sorted_hist.size(1)-1)*upper_threshold)][:,None,None]
mask = ((hist-lower)/(upper-lower)).clip(max=1, min=0)
return (mask.cpu(),)
NODE_CLASS_MAPPINGS = {
"DynamicSampler": DynamicSampler,
"MeasuredSampler": MeasuredSamplerNode,
"ResolveMaskPromise": ResolveMaskPromise,
}
NODE_DISPLAY_NAME_MAPPINGS = {}