(self, batch: Dict)
| 135 | return loss_mean, loss_dict |
| 136 | |
| 137 | def shared_step(self, batch: Dict) -> Any: |
| 138 | x = self.get_input(batch) |
| 139 | if self.lr_scale is not None: |
| 140 | lr_x = F.interpolate(x, scale_factor=1 / self.lr_scale, mode="bilinear", align_corners=False) |
| 141 | lr_x = F.interpolate(lr_x, scale_factor=self.lr_scale, mode="bilinear", align_corners=False) |
| 142 | lr_z = self.encode_first_stage(lr_x, batch) |
| 143 | batch["lr_input"] = lr_z |
| 144 | |
| 145 | x = x.permute(0, 2, 1, 3, 4).contiguous() # (B, T, C, H, W) -> (B, C, T, H, W) |
| 146 | hq_video = x # (B, C, T, H, W) |
| 147 | x = self.encode_first_stage(x, batch) |
| 148 | x = x.permute(0, 2, 1, 3, 4).contiguous() # (B, C, T, H, W) -> (B, T, C, H, W) |
| 149 | |
| 150 | if 'lq' in batch.keys(): |
| 151 | # print('LQ is NOT None') |
| 152 | lq = batch['lq'].to(self.dtype) |
| 153 | lq = lq.permute(0, 2, 1, 3, 4).contiguous() |
| 154 | lq = self.encode_first_stage(lq, batch) |
| 155 | lq = lq.permute(0, 2, 1, 3, 4).contiguous() |
| 156 | batch['lq'] = lq |
| 157 | |
| 158 | # Uncomment for t2v training, |
| 159 | # batch['lq'] = None |
| 160 | |
| 161 | gc.collect() |
| 162 | torch.cuda.empty_cache() |
| 163 | loss, loss_dict = self(x, hq_video, batch) |
| 164 | return loss, loss_dict |
| 165 | |
| 166 | def get_input(self, batch): |
| 167 | return batch[self.input_key].to(self.dtype) |
nothing calls this directly
no test coverage detected