1 Commits
Author SHA1 Message Date
Austin Mroz de6dea5425 Oversample in second pass
In an attempt to improve visual quality, the oversampling is instead
performed as a second pass at the end. The hope was that this would
combine the promising selection code with an existent and known
working solution for improving the visual fidelity of the subregion, but
performance is still worse than expected.
2024-05-30 15:12:33 -05:00
2 changed files with 75 additions and 141 deletions
+1
View File
@@ -0,0 +1 @@
__pycache__/
+74 -141
View File
@@ -8,6 +8,24 @@ from comfy.k_diffusion import sampling
from comfy.samplers import KSAMPLER
from tqdm.auto import trange
from comfy import latent_formats
latent_factors = torch.tensor(latent_formats.SD15().latent_rgb_factors)
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")
@@ -18,12 +36,13 @@ def dump_image(tensor):
if len(tensor.shape) == 3:
tensor = tensor.unsqueeze(1).expand(-1,3,-1,-1)
tensor = tensor.permute(0,2,3,1)[0]
tensor @= latent_factors
h,w,c = tensor.shape
b = np.clip(tensor.numpy() * 255, 0, 255).astype(np.uint8)
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., sub_weight=3, early_terminate=False, resolution_mult=1.25):
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.):
"""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]])
@@ -31,11 +50,13 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
hist = torch.zeros((n,h,w), device=x.device, dtype=x.dtype)
kx, ky = torch.meshgrid(torch.linspace(-h/2,h/2,h),torch.linspace(-w/2,w/2,w), indexing="ij")
base_mask = (kx*kx + ky*ky).sqrt().to(x.device)
kx, ky = torch.meshgrid(torch.linspace(-h/4,h/4,h//2),torch.linspace(-w/4,w/4,w//2), indexing="ij")
base_submask = (kx*kx + ky*ky).sqrt().to(x.device)
cw = w-w//4+1
ch = h-h//4+1
kernel = torch.ones((1,1,h//4,w//4), dtype=x.dtype, device=x.device)
prev_denoised = None
oversample_performed = False
early_terminate = 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)
@@ -43,57 +64,15 @@ 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
#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 = []
subinds = []
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:
oversample_performed = True
subx = torch.stack(sublocs)
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
subdn = model(subx, sigma_hat*s_in*math.sqrt(2), **extra_args)
#downscale
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]
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, focalm)
prev_denoised = denoised
d = sampling.to_d(x, sigma_hat, denoised)
@@ -102,118 +81,72 @@ 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 and oversample_performed:
if early_terminate:
break
#Debugging code to dump the measured distances
#summed = torch.nn.functional.conv2d(hist, kernel, groups=1)
#dump_image(hist)
summed = torch.nn.functional.conv2d(hist, kernel, groups=1)
focal = torch.unravel_index(summed.view((n,ch*cw)).argmax(1), (ch,cw))
focal = torch.stack(focal).transpose(0,1)
sublocs = []
subinds = []
focalm = []
for j,ind in enumerate(focal):
focalm.append(summed[j, *ind])
if summed[j, *ind] > 1000:
subx = x[j,:,ind[0]:ind[0]+h//4,ind[1]:ind[1]+w//4]
sublocs.append(subx)
subinds.append(j)
if len(sublocs) > 0:
print("doing subrender")
subx = torch.stack(sublocs)
subx = interpolate(subx, size=(h//2,w//2), mode='bicubic')
subnoise = model.noise[:,:,:h//2,:w//2]
subnoise = model.inner_model.inner_model.model_sampling.noise_scaling(sigmas[len(sigmas)//2], subnoise, subx, 0)
subres = sampling.sample_euler(model, subx, sigmas[len(sigmas)//2:], extra_args, callback, disable, s_churn, s_tmin, s_tmax, s_noise)
#downscale
subres = interpolate(subres, size=(h//4,w//4))
#TODO: Does added noise need to be filtered out?
for j,dn in zip(subinds,subres):
subx = x[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)
return x
class DynamicSampler:
@classmethod
def INPUT_TYPES(s):
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 {"required": {}}
RETURN_TYPES = ("SAMPLER",)
CATEGORY = "sampling/custom_sampling/samplers"
FUNCTION = "get_sampler"
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)),)
def get_sampler(self):
return (KSAMPLER(sample_dynamic),)
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:
class LatentAsImage:
"""Converts a latent to an image with minimal processing. Provides a means of viewing
latent channels visually with minimal overhead"""
@classmethod
def INPUT_TYPES(s):
return {"required": {"sampler": ("SAMPLER",),}}
RETURN_TYPES = ("SAMPLER", "MASK_PROMISE")
CATEGORY = "sampling/custom_sampling/samplers"
return {"required": {
"latent" : ("LATENT",),
}}
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(),)
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_(),)
NODE_CLASS_MAPPINGS = {
"DynamicSampler": DynamicSampler,
"MeasuredSampler": MeasuredSamplerNode,
"ResolveMaskPromise": ResolveMaskPromise,
}
NODE_DISPLAY_NAME_MAPPINGS = {}