(self, model_input, tile_size, tile_stride, tile_device, tile_dtype)
| 18 | |
| 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.shape |
| 23 | 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_input |
| 32 | |
| 33 | |
| 34 | def tiled_inference(self, forward_fn, model_input, tile_batch_size, inference_device, inference_dtype, tile_device, tile_dtype): |
no test coverage detected