MCPcopy Create free account
hub / github.com/ali-vilab/ACE_plus / preprocess

Method preprocess

inference/utils.py:46–122  ·  view source on GitHub ↗
(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)

Source from the content-addressed store, hash-verified

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)

Callers 2

__call__Method · 0.45
__call__Method · 0.45

Calls 1

image_checkMethod · 0.95

Tested by

no test coverage detected