(self,
reference_image=None,
edit_image=None,
edit_mask=None,
height=1024,
width=1024,
repainting_scale = 1.0,
keep_pixels = False,
keep_pixels_rate = 0.8,
use_change = False)
| 44 | |
| 45 | |
| 46 | def preprocess(self, |
| 47 | reference_image=None, |
| 48 | edit_image=None, |
| 49 | edit_mask=None, |
| 50 | height=1024, |
| 51 | width=1024, |
| 52 | repainting_scale = 1.0, |
| 53 | keep_pixels = False, |
| 54 | keep_pixels_rate = 0.8, |
| 55 | use_change = False): |
| 56 | reference_image = self.image_check(reference_image) |
| 57 | edit_image = self.image_check(edit_image) |
| 58 | # for reference generation |
| 59 | if edit_image is None: |
| 60 | edit_image = torch.zeros([3, height, width]) |
| 61 | edit_mask = torch.ones([1, height, width]) |
| 62 | else: |
| 63 | if edit_mask is None: |
| 64 | _, eH, eW = edit_image.shape |
| 65 | edit_mask = np.ones((eH, eW)) |
| 66 | else: |
| 67 | edit_mask = np.asarray(edit_mask) |
| 68 | edit_mask = np.where(edit_mask > 128, 1, 0) |
| 69 | edit_mask = edit_mask.astype( |
| 70 | np.float32) if np.any(edit_mask) else np.ones_like(edit_mask).astype( |
| 71 | np.float32) |
| 72 | edit_mask = torch.tensor(edit_mask).unsqueeze(0) |
| 73 | |
| 74 | edit_image = edit_image * (1 - edit_mask * repainting_scale) |
| 75 | |
| 76 | |
| 77 | out_h, out_w = edit_image.shape[-2:] |
| 78 | |
| 79 | assert edit_mask is not None |
| 80 | if reference_image is not None: |
| 81 | _, H, W = reference_image.shape |
| 82 | _, eH, eW = edit_image.shape |
| 83 | if not keep_pixels: |
| 84 | # align height with edit_image |
| 85 | scale = eH / H |
| 86 | tH, tW = eH, int(W * scale) |
| 87 | reference_image = T.Resize((tH, tW), interpolation=T.InterpolationMode.BILINEAR, antialias=True)( |
| 88 | reference_image) |
| 89 | else: |
| 90 | # padding |
| 91 | if H >= keep_pixels_rate * eH: |
| 92 | tH = int(eH * keep_pixels_rate) |
| 93 | scale = tH/H |
| 94 | tW = int(W * scale) |
| 95 | reference_image = T.Resize((tH, tW), interpolation=T.InterpolationMode.BILINEAR, antialias=True)( |
| 96 | reference_image) |
| 97 | rH, rW = reference_image.shape[-2:] |
| 98 | delta_w = 0 |
| 99 | delta_h = eH - rH |
| 100 | padding = (delta_w // 2, delta_h // 2, delta_w - (delta_w // 2), delta_h - (delta_h // 2)) |
| 101 | reference_image = T.Pad(padding, fill=0, padding_mode="constant")(reference_image) |
| 102 | edit_image = torch.cat([reference_image, edit_image], dim=-1) |
| 103 | edit_mask = torch.cat([torch.zeros([1, reference_image.shape[1], reference_image.shape[2]]), edit_mask], dim=-1) |
no test coverage detected