modelscope/DiffSynth-Painter
14
1import torch2from einops import rearrange, repeat3 4 5class TileWorker:6 def __init__(self):7 pass8 9 10 def mask(self, height, width, border_width):11 # Create a mask with shape (height, width).12 # The centre area is filled with 1, and the border line is filled with values in range (0, 1].13 x = torch.arange(height).repeat(width, 1).T14 y = torch.arange(width).repeat(height, 1)15 mask = torch.stack([x + 1, height - x, y + 1, width - y]).min(dim=0).values16 mask = (mask / border_width).clip(0, 1)17 return mask18 19 20 def tile(self, model_input, tile_size, tile_stride, tile_device, tile_dtype):21 # Convert a tensor (b, c, h, w) to (b, c, tile_size, tile_size, tile_num)22 batch_size, channel, _, _ = model_input.shape23 model_input = model_input.to(device=tile_device, dtype=tile_dtype)24 unfold_operator = torch.nn.Unfold(25 kernel_size=(tile_size, tile_size),26 stride=(tile_stride, tile_stride)27 )28 model_input = unfold_operator(model_input)29 model_input = model_input.view((batch_size, channel, tile_size, tile_size, -1))30 31 return model_input32 33 34 def tiled_inference(self, forward_fn, model_input, tile_batch_size, inference_device, inference_dtype, tile_device, tile_dtype):35 # Call y=forward_fn(x) for each tile36 tile_num = model_input.shape[-1]37 model_output_stack = []38 39 for tile_id in range(0, tile_num, tile_batch_size):40 41 # process input42 tile_id_ = min(tile_id + tile_batch_size, tile_num)43 x = model_input[:, :, :, :, tile_id: tile_id_]44 x = x.to(device=inference_device, dtype=inference_dtype)45 x = rearrange(x, "b c h w n -> (n b) c h w")46 47 # process output48 y = forward_fn(x)49 y = rearrange(y, "(n b) c h w -> b c h w n", n=tile_id_-tile_id)50 y = y.to(device=tile_device, dtype=tile_dtype)51 model_output_stack.append(y)52 53 model_output = torch.concat(model_output_stack, dim=-1)54 return model_output55 56 57 def io_scale(self, model_output, tile_size):58 # Determine the size modification happend in forward_fn59 # We only consider the same scale on height and width.60 io_scale = model_output.shape[2] / tile_size61 return io_scale62 63 64 def untile(self, model_output, height, width, tile_size, tile_stride, border_width, tile_device, tile_dtype):65 # The reversed function of tile66 mask = self.mask(tile_size, tile_size, border_width)67 mask = mask.to(device=tile_device, dtype=tile_dtype)68 mask = rearrange(mask, "h w -> 1 1 h w 1")69 model_output = model_output * mask70 71 fold_operator = torch.nn.Fold(72 output_size=(height, width),73 kernel_size=(tile_size, tile_size),74 stride=(tile_stride, tile_stride)75 )76 mask = repeat(mask[0, 0, :, :, 0], "h w -> 1 (h w) n", n=model_output.shape[-1])77 model_output = rearrange(model_output, "b c h w n -> b (c h w) n")78 model_output = fold_operator(model_output) / fold_operator(mask)79 80 return model_output81 82 83 def tiled_forward(self, forward_fn, model_input, tile_size, tile_stride, tile_batch_size=1, tile_device="cpu", tile_dtype=torch.float32, border_width=None):84 # Prepare85 inference_device, inference_dtype = model_input.device, model_input.dtype86 height, width = model_input.shape[2], model_input.shape[3]87 border_width = int(tile_stride*0.5) if border_width is None else border_width88 89 # tile90 model_input = self.tile(model_input, tile_size, tile_stride, tile_device, tile_dtype)91 92 # inference93 model_output = self.tiled_inference(forward_fn, model_input, tile_batch_size, inference_device, inference_dtype, tile_device, tile_dtype)94 95 # resize96 io_scale = self.io_scale(model_output, tile_size)97 height, width = int(height*io_scale), int(width*io_scale)98 tile_size, tile_stride = int(tile_size*io_scale), int(tile_stride*io_scale)99 border_width = int(border_width*io_scale)100 101 # untile102 model_output = self.untile(model_output, height, width, tile_size, tile_stride, border_width, tile_device, tile_dtype)103 104 # Done!105 model_output = model_output.to(device=inference_device, dtype=inference_dtype)106 return model_output