File size: 4,383 Bytes
38d5abe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107

import random

import numpy as np
from PIL import Image, ImageEnhance, ImageFilter



def binarize_mask(mask, threshold=0):
    return mask.point(lambda p: 255 if p > threshold else 0).convert("L")


def square_crop_by_mask(image, mask, scale=1.2):
    W, H = image.size
    bbox = mask.getbbox()
    if bbox is None:
        cx, cy = W / 2.0, H / 2.0
        side = float(min(W, H))
    else:
        l, u, r, lo = bbox
        cx = (l + r) / 2.0
        cy = (u + lo) / 2.0
        side = max(r - l, lo - u) * scale
    side_i = max(1, int(round(side)))
    left = int(round(cx - side_i / 2.0))
    top = int(round(cy - side_i / 2.0))
    box = (left, top, left + side_i, top + side_i)
    # PIL fills out-of-bounds crop regions with 0: black for the RGB image and
    # "unmasked" (0) for the L mask - exactly the black padding we want.
    return image.crop(box), mask.crop(box), box



def process_source(source_image, source_mask, image_size=1024, crop_scale=1.2, binarize_threshold=0):
    source_mask = binarize_mask(source_mask, binarize_threshold)
    source_image_c, source_mask_c, crop_box = square_crop_by_mask(source_image, source_mask, crop_scale)
    size = (image_size, image_size)
    image = source_image_c.resize(size, Image.LANCZOS)        # GT crop
    mask = source_mask_c.resize(size, Image.NEAREST)
    white = Image.new("RGB", size, (255, 255, 255))
    # composite(white, image, mask): mask==255 -> white, mask==0 -> image
    background_image = Image.composite(white, image, mask)
    return background_image, image, crop_box, source_mask_c


def process_reference(ref_image, ref_mask, image_size=1024, crop_scale=1.2,
                      binarize_threshold=0, augment=False, p_aug=0.8):
    ref_mask = binarize_mask(ref_mask, binarize_threshold)
    ref_image_c, ref_mask_c, _ = square_crop_by_mask(ref_image, ref_mask, crop_scale)
    size = (image_size, image_size)
    ref_image_r = ref_image_c.resize(size, Image.LANCZOS)
    ref_mask_r = ref_mask_c.resize(size, Image.NEAREST)
    white = Image.new("RGB", size, (255, 255, 255))
    # composite(ref_image, white, mask): mask==255 -> object, mask==0 -> white
    ref_image_out = Image.composite(ref_image_r, white, ref_mask_r)
    if augment:
        ref_image_out = augment_ref_image(ref_image_out, ref_mask_r, p_aug=p_aug)
    return ref_image_out


def paste_back(generated_crop, source_image, crop_box, source_mask_cropped, feather=0):
    left, top, right, bottom = crop_box
    crop_side = right - left
    generated_crop = generated_crop.resize((crop_side, crop_side), Image.LANCZOS)
    original_crop = source_image.crop(crop_box)
    if feather > 0:
        source_mask_cropped = source_mask_cropped.filter(ImageFilter.GaussianBlur(feather))
    edited_crop = Image.composite(generated_crop, original_crop, source_mask_cropped)
    final_image = source_image.copy()
    final_image.paste(edited_crop, (left, top))
    return final_image


def _adjust_hue(img, shift):
    """Shift the hue of an RGB image by ``shift`` (in 0..255 units)."""
    hsv = img.convert("HSV")
    h, s, v = hsv.split()
    h_arr = (np.asarray(h).astype(np.int16) + int(round(shift))) % 256
    h = Image.fromarray(h_arr.astype(np.uint8), "L")
    return Image.merge("HSV", (h, s, v)).convert("RGB")


def augment_ref_image(ref_image, mask, p_aug=0.8):
    img = ref_image.convert("RGB").copy()
    m = mask.convert("L")

    # Horizontal flip is safe: the object sits on a uniform white background.
    if random.random() < 0.5:
        img = img.transpose(Image.FLIP_LEFT_RIGHT)
        m = m.transpose(Image.FLIP_LEFT_RIGHT)

    # Colour jitter on the whole image; background is re-whitened below.
    if random.random() < p_aug:
        img = ImageEnhance.Brightness(img).enhance(random.uniform(0.85, 1.15))
    if random.random() < p_aug:
        img = ImageEnhance.Contrast(img).enhance(random.uniform(0.85, 1.15))
    if random.random() < p_aug:
        img = ImageEnhance.Color(img).enhance(random.uniform(0.85, 1.15))
    if random.random() < p_aug:
        img = ImageEnhance.Sharpness(img).enhance(random.uniform(0.9, 1.1))
    if random.random() < p_aug:
        img = _adjust_hue(img, random.uniform(-0.06, 0.06) * 255.0)

    # Keep the (jittered) object, reset everything outside the mask to white.
    white = Image.new("RGB", img.size, (255, 255, 255))
    return Image.composite(img, white, m)