(self, rendered_frames, **kwargs)
| 53 | |
| 54 | @torch.no_grad() |
| 55 | def __call__(self, rendered_frames, **kwargs): |
| 56 | # Preprocess |
| 57 | processed_images = self.process_images(rendered_frames) |
| 58 | |
| 59 | # Input |
| 60 | input_tensor = torch.cat((processed_images[:-2], processed_images[2:]), dim=1) |
| 61 | |
| 62 | # Interpolate |
| 63 | output_tensor = self.process_tensors(input_tensor, scale=self.scale, batch_size=self.batch_size) |
| 64 | |
| 65 | if self.interpolate: |
| 66 | # Blend |
| 67 | input_tensor = torch.cat((processed_images[1:-1], output_tensor), dim=1) |
| 68 | output_tensor = self.process_tensors(input_tensor, scale=self.scale, batch_size=self.batch_size) |
| 69 | processed_images[1:-1] = output_tensor |
| 70 | else: |
| 71 | processed_images[1:-1] = (processed_images[1:-1] + output_tensor) / 2 |
| 72 | |
| 73 | # To images |
| 74 | output_images = self.decode_images(processed_images) |
| 75 | if output_images[0].size != rendered_frames[0].size: |
| 76 | output_images = [image.resize(rendered_frames[0].size) for image in output_images] |
| 77 | return output_images |
nothing calls this directly
no test coverage detected