(self, frames, depths, cond_times=None)
| 581 | param.requires_grad = False |
| 582 | |
| 583 | def forward_encoder(self, frames, depths, cond_times=None): |
| 584 | frames = torch.cat([frames, depths], dim=2) |
| 585 | encoder_output = self.gs_predictor.encoder(frames, timestamp=None) |
| 586 | |
| 587 | return encoder_output |
| 588 | |
| 589 | def upsampling(self, output): |
| 590 | input_views = output.shape[1] |