| 113 | |
| 114 | |
| 115 | def build_mask(self, data, is_bound): |
| 116 | _, _, H, W = data.shape |
| 117 | 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) // 4 |
| 120 | pad = torch.ones_like(h) * border_width |
| 121 | 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 - w |
| 126 | ]).min(dim=0).values |
| 127 | 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 mask |
| 131 | |
| 132 | |
| 133 | def tiled_forward(self, forward_fn, model_input, tile_size, tile_stride, tile_device="cpu", tile_dtype=torch.float32, border_width=None): |