MCPcopy Create free account
hub / github.com/NJU-PCALab/STAR / shared_step

Method shared_step

cogvideox-based/sat/diffusion_video.py:137–164  ·  view source on GitHub ↗
(self, batch: Dict)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 3

get_inputMethod · 0.95
encode_first_stageMethod · 0.95
toMethod · 0.80

Tested by

no test coverage detected