fix many inpaint bugs (#731)

fix many inpaint bugs
This commit is contained in:
lllyasviel
2023-10-18 06:22:08 -07:00
committed by GitHub
parent d2c8f16082
commit 9660daff94
6 changed files with 131 additions and 96 deletions
+80 -61
View File
@@ -43,6 +43,12 @@ def morphological_open(x):
return x_int32.clip(0, 255).astype(np.uint8)
def up255(x, t=0):
y = np.zeros_like(x).astype(np.uint8)
y[x > t] = 255
return y
def imsave(x, path):
x = Image.fromarray(x)
x.save(path)
@@ -75,21 +81,25 @@ def compute_initial_abcd(x):
b = np.max(indices[0]) + 65
c = np.min(indices[1]) - 64
d = np.max(indices[1]) + 65
abp = (b + a) // 2
abm = (b - a) // 2
cdp = (d + c) // 2
cdm = (d - c) // 2
l = max(abm, cdm)
a = abp - l
b = abp + l
c = cdp - l
d = cdp + l
a, b, c, d = regulate_abcd(x, a, b, c, d)
return a, b, c, d
def area_abcd(a, b, c, d):
return (b - a) * (d - c)
def solve_abcd(x, a, b, c, d, outpaint):
H, W = x.shape[:2]
if outpaint:
return 0, H, 0, W
min_area = (min(H, W) ** 2) * 0.5
while True:
if area_abcd(a, b, c, d) >= min_area:
if b - a > H * 0.618 and d - c > W * 0.618:
break
add_h = (b - a) < (d - c)
@@ -119,7 +129,7 @@ def fooocus_fill(image, mask):
area = np.where(mask < 127)
store = raw_image[area]
for k, repeats in [(64, 4), (32, 4), (16, 4), (4, 4), (2, 4)]:
for k, repeats in [(512, 2), (256, 2), (128, 4), (64, 4), (33, 8), (15, 8), (5, 16), (3, 16)]:
for _ in range(repeats):
current_image = box_blur(current_image, k)
current_image[area] = store
@@ -129,98 +139,107 @@ def fooocus_fill(image, mask):
class InpaintWorker:
def __init__(self, image, mask, is_outpaint):
# mask processing
self.mask_raw_soft = morphological_open(mask)
self.mask_raw_fg = (self.mask_raw_soft == 255).astype(np.uint8) * 255
self.mask_raw_bg = (self.mask_raw_soft == 0).astype(np.uint8) * 255
self.mask_raw_trim = 255 - np.maximum(self.mask_raw_fg, self.mask_raw_bg)
# image processing
self.image_raw = fooocus_fill(image, self.mask_raw_fg)
# log all images
# imsave(self.image_raw, 'image_raw.png')
# imsave(self.mask_raw_soft, 'mask_raw_soft.png')
# imsave(self.mask_raw_fg, 'mask_raw_fg.png')
# imsave(self.mask_raw_bg, 'mask_raw_bg.png')
# imsave(self.mask_raw_trim, 'mask_raw_trim.png')
# compute abcd
a, b, c, d = compute_initial_abcd(self.mask_raw_bg < 127)
a, b, c, d = solve_abcd(self.mask_raw_bg, a, b, c, d, outpaint=is_outpaint)
a, b, c, d = compute_initial_abcd(mask > 0)
a, b, c, d = solve_abcd(mask, a, b, c, d, outpaint=is_outpaint)
# interested area
self.interested_area = (a, b, c, d)
self.mask_interested_soft = self.mask_raw_soft[a:b, c:d]
self.mask_interested_fg = self.mask_raw_fg[a:b, c:d]
self.mask_interested_bg = self.mask_raw_bg[a:b, c:d]
self.mask_interested_trim = self.mask_raw_trim[a:b, c:d]
self.image_interested = self.image_raw[a:b, c:d]
self.interested_mask = mask[a:b, c:d]
self.interested_image = image[a:b, c:d]
# resize to make images ready for diffusion
H, W, C = self.image_interested.shape
k = (1024.0 ** 2.0 / float(H * W)) ** 0.5
H, W, C = self.interested_image.shape
k = ((1024.0 ** 2.0) / float(H * W)) ** 0.5
H = int(np.ceil(float(H) * k / 16.0)) * 16
W = int(np.ceil(float(W) * k / 16.0)) * 16
self.image_ready = resample_image(self.image_interested, W, H)
self.mask_ready = resample_image(self.mask_interested_soft, W, H)
self.interested_mask = up255(resample_image(self.interested_mask, W, H), t=127)
self.interested_image = resample_image(self.interested_image, W, H)
self.interested_fill = fooocus_fill(self.interested_image, self.interested_mask)
# soft pixels
self.mask = morphological_open(mask)
self.image = image
# ending
self.latent = None
self.latent_after_swap = None
self.swapped = False
self.latent_mask = None
self.inpaint_head_feature = None
return
def load_inpaint_guidance(self, latent, mask, model_path):
def load_latent(self,
latent_fill,
latent_inpaint,
latent_mask,
latent_swap=None,
inpaint_head_model_path=None):
global inpaint_head
assert inpaint_head_model_path is not None
self.latent = latent_fill
self.latent_mask = latent_mask
self.latent_after_swap = latent_swap
if inpaint_head is None:
inpaint_head = InpaintHead()
sd = torch.load(model_path, map_location='cpu')
sd = torch.load(inpaint_head_model_path, map_location='cpu')
inpaint_head.load_state_dict(sd)
process_latent_in = pipeline.xl_base_patched.unet.model.process_latent_in
latent = process_latent_in(latent)
B, C, H, W = latent.shape
mask = torch.nn.functional.interpolate(mask, size=(H, W), mode="bilinear")
mask = mask.round()
feed = torch.cat([mask, latent], dim=1)
feed = torch.cat([
latent_mask,
pipeline.xl_base_patched.unet.model.process_latent_in(latent_inpaint)
], dim=1)
inpaint_head.to(device=feed.device, dtype=feed.dtype)
self.inpaint_head_feature = inpaint_head(feed)
return
def load_latent(self, latent, mask, latent_after_swap=None):
self.latent = latent
self.latent_mask = mask
self.latent_after_swap = latent_after_swap
def swap(self):
if self.latent_after_swap is not None:
self.latent, self.latent_after_swap = self.latent_after_swap, self.latent
if self.swapped:
return
if self.latent is None:
return
if self.latent_after_swap is None:
return
self.latent, self.latent_after_swap = self.latent_after_swap, self.latent
self.swapped = True
return
def unswap(self):
if not self.swapped:
return
if self.latent is None:
return
if self.latent_after_swap is None:
return
self.latent, self.latent_after_swap = self.latent_after_swap, self.latent
self.swapped = False
return
def color_correction(self, img):
fg = img.astype(np.float32)
bg = self.image_raw.copy().astype(np.float32)
w = self.mask_raw_soft[:, :, None].astype(np.float32) / 255.0
bg = self.image.copy().astype(np.float32)
w = self.mask[:, :, None].astype(np.float32) / 255.0
y = fg * w + bg * (1 - w)
return y.clip(0, 255).astype(np.uint8)
def post_process(self, img):
a, b, c, d = self.interested_area
content = resample_image(img, d - c, b - a)
result = self.image_raw.copy()
result = self.image.copy()
result[a:b, c:d] = content
result = self.color_correction(result)
return result
def visualize_mask_processing(self):
result = self.image_raw // 4
a, b, c, d = self.interested_area
result[a:b, c:d] += 64
result[self.mask_raw_trim > 127] += 64
result[self.mask_raw_fg > 127] += 128
return [result, self.mask_raw_soft, self.image_ready, self.mask_ready]
return [self.interested_fill, self.interested_mask, self.image, self.mask]