Team Ai
Apppublic

hugging-apps/echo-memory

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
tiler.py234 linesDownload Raw Back to models
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 happened 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_output107    108 109 110class FastTileWorker:111    def __init__(self):112        pass113 114 115    def build_mask(self, data, is_bound):116        _, _, H, W = data.shape117        h = repeat(torch.arange(H), "H -> H W", H=H, W=W)118        w = repeat(torch.arange(W), "W -> H W", H=H, W=W)119        border_width = (H + W) // 4120        pad = torch.ones_like(h) * border_width121        mask = torch.stack([122            pad if is_bound[0] else h + 1,123            pad if is_bound[1] else H - h,124            pad if is_bound[2] else w + 1,125            pad if is_bound[3] else W - w126        ]).min(dim=0).values127        mask = mask.clip(1, border_width)128        mask = (mask / border_width).to(dtype=data.dtype, device=data.device)129        mask = rearrange(mask, "H W -> 1 H W")130        return mask131 132 133    def tiled_forward(self, forward_fn, model_input, tile_size, tile_stride, tile_device="cpu", tile_dtype=torch.float32, border_width=None):134        # Prepare135        B, C, H, W = model_input.shape136        border_width = int(tile_stride*0.5) if border_width is None else border_width137        weight = torch.zeros((1, 1, H, W), dtype=tile_dtype, device=tile_device)138        values = torch.zeros((B, C, H, W), dtype=tile_dtype, device=tile_device)139 140        # Split tasks141        tasks = []142        for h in range(0, H, tile_stride):143            for w in range(0, W, tile_stride):144                if (h-tile_stride >= 0 and h-tile_stride+tile_size >= H) or (w-tile_stride >= 0 and w-tile_stride+tile_size >= W):145                    continue146                h_, w_ = h + tile_size, w + tile_size147                if h_ > H: h, h_ = H - tile_size, H148                if w_ > W: w, w_ = W - tile_size, W149                tasks.append((h, h_, w, w_))150        151        # Run152        for hl, hr, wl, wr in tasks:153            # Forward154            hidden_states_batch = forward_fn(hl, hr, wl, wr).to(dtype=tile_dtype, device=tile_device)155 156            mask = self.build_mask(hidden_states_batch, is_bound=(hl==0, hr>=H, wl==0, wr>=W))157            values[:, :, hl:hr, wl:wr] += hidden_states_batch * mask158            weight[:, :, hl:hr, wl:wr] += mask159        values /= weight160        return values161 162 163 164class TileWorker2Dto3D:165    """166    Process 3D tensors, but only enable TileWorker on 2D.167    """168    def __init__(self):169        pass170 171 172    def build_mask(self, T, H, W, dtype, device, is_bound, border_width):173        t = repeat(torch.arange(T), "T -> T H W", T=T, H=H, W=W)174        h = repeat(torch.arange(H), "H -> T H W", T=T, H=H, W=W)175        w = repeat(torch.arange(W), "W -> T H W", T=T, H=H, W=W)176        border_width = (H + W) // 4 if border_width is None else border_width177        pad = torch.ones_like(h) * border_width178        mask = torch.stack([179            pad if is_bound[0] else t + 1,180            pad if is_bound[1] else T - t,181            pad if is_bound[2] else h + 1,182            pad if is_bound[3] else H - h,183            pad if is_bound[4] else w + 1,184            pad if is_bound[5] else W - w185        ]).min(dim=0).values186        mask = mask.clip(1, border_width)187        mask = (mask / border_width).to(dtype=dtype, device=device)188        mask = rearrange(mask, "T H W -> 1 1 T H W")189        return mask190 191 192    def tiled_forward(193        self,194        forward_fn,195        model_input,196        tile_size, tile_stride,197        tile_device="cpu", tile_dtype=torch.float32,198        computation_device="cuda", computation_dtype=torch.float32,199        border_width=None, scales=[1, 1, 1, 1],200        progress_bar=lambda x:x201    ):202        B, C, T, H, W = model_input.shape203        scale_C, scale_T, scale_H, scale_W = scales204        tile_size_H, tile_size_W = tile_size205        tile_stride_H, tile_stride_W = tile_stride206 207        value = torch.zeros((B, int(C*scale_C), int(T*scale_T), int(H*scale_H), int(W*scale_W)), dtype=tile_dtype, device=tile_device)208        weight = torch.zeros((1, 1, int(T*scale_T), int(H*scale_H), int(W*scale_W)), dtype=tile_dtype, device=tile_device)209 210        # Split tasks211        tasks = []212        for h in range(0, H, tile_stride_H):213            for w in range(0, W, tile_stride_W):214                if (h-tile_stride_H >= 0 and h-tile_stride_H+tile_size_H >= H) or (w-tile_stride_W >= 0 and w-tile_stride_W+tile_size_W >= W):215                    continue216                h_, w_ = h + tile_size_H, w + tile_size_W217                if h_ > H: h, h_ = max(H - tile_size_H, 0), H218                if w_ > W: w, w_ = max(W - tile_size_W, 0), W219                tasks.append((h, h_, w, w_))220 221        # Run222        for hl, hr, wl, wr in progress_bar(tasks):223            mask = self.build_mask(224                int(T*scale_T), int((hr-hl)*scale_H), int((wr-wl)*scale_W),225                tile_dtype, tile_device,226                is_bound=(True, True, hl==0, hr>=H, wl==0, wr>=W),227                border_width=border_width228            )229            grid_input = model_input[:, :, :, hl:hr, wl:wr].to(dtype=computation_dtype, device=computation_device)230            grid_output = forward_fn(grid_input).to(dtype=tile_dtype, device=tile_device)231            value[:, :, :, int(hl*scale_H):int(hr*scale_H), int(wl*scale_W):int(wr*scale_W)] += grid_output * mask232            weight[:, :, :, int(hl*scale_H):int(hr*scale_H), int(wl*scale_W):int(wr*scale_W)] += mask233        value = value / weight234        return value