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.
This commit is contained in:
@@ -0,0 +1 @@
|
||||
__pycache__/
|
||||
@@ -8,6 +8,9 @@ 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
|
||||
@@ -33,6 +36,7 @@ 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")
|
||||
@@ -46,6 +50,8 @@ 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)
|
||||
@@ -65,44 +71,8 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
|
||||
filt = torch.fft.ifft2(mask*torch.fft.fft2(denoised)).real
|
||||
if prev_denoised is not None:
|
||||
ext = prev_denoised+filt-denoised
|
||||
dist = (ext*ext).sum(1)/-dt
|
||||
dist = (ext*ext).sum(1)/math.sqrt(-dt)
|
||||
hist+= dist
|
||||
summed = torch.nn.functional.conv2d(dist, 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]
|
||||
#scale+noise subx
|
||||
sublocs.append(subx)
|
||||
subinds.append(j)
|
||||
if len(sublocs) > 0:
|
||||
early_terminate = 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')
|
||||
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?
|
||||
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)
|
||||
prev_denoised = denoised
|
||||
d = sampling.to_d(x, sigma_hat, denoised)
|
||||
|
||||
@@ -114,7 +84,39 @@ def sample_dynamic(model, x, sigmas, extra_args=None, callback=None, disable=Non
|
||||
if early_terminate:
|
||||
break
|
||||
summed = torch.nn.functional.conv2d(hist, kernel, groups=1)
|
||||
dump_image(hist)
|
||||
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:
|
||||
|
||||
Reference in New Issue
Block a user