(self, videos, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16))
| 786 | |
| 787 | |
| 788 | def encode(self, videos, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)): |
| 789 | |
| 790 | videos = [video.to("cpu") for video in videos] |
| 791 | hidden_states = [] |
| 792 | for video in videos: |
| 793 | video = video.unsqueeze(0) |
| 794 | if tiled: |
| 795 | tile_size = (tile_size[0] * 8, tile_size[1] * 8) |
| 796 | tile_stride = (tile_stride[0] * 8, tile_stride[1] * 8) |
| 797 | hidden_state = self.tiled_encode(video, device, tile_size, tile_stride) |
| 798 | else: |
| 799 | hidden_state = self.single_encode(video, device) |
| 800 | hidden_state = hidden_state.squeeze(0) |
| 801 | hidden_states.append(hidden_state) |
| 802 | hidden_states = torch.stack(hidden_states) |
| 803 | return hidden_states |
| 804 | |
| 805 | |
| 806 | def decode(self, hidden_states, device, tiled=False, tile_size=(34, 34), tile_stride=(18, 16)): |
nothing calls this directly
no test coverage detected