(self, sample, tile_size=64, tile_stride=32)
| 43 | self.conv_out = torch.nn.Conv2d(128, 3, kernel_size=3, padding=1) |
| 44 | |
| 45 | def tiled_forward(self, sample, tile_size=64, tile_stride=32): |
| 46 | hidden_states = TileWorker().tiled_forward( |
| 47 | lambda x: self.forward(x), |
| 48 | sample, |
| 49 | tile_size, |
| 50 | tile_stride, |
| 51 | tile_device=sample.device, |
| 52 | tile_dtype=sample.dtype |
| 53 | ) |
| 54 | return hidden_states |
| 55 | |
| 56 | def forward(self, sample, tiled=False, tile_size=64, tile_stride=32, **kwargs): |
| 57 | # For VAE Decoder, we do not need to apply the tiler on each layer. |
no test coverage detected