BonanDing commited on
Commit
e8059ff
·
1 Parent(s): 448fb39

Thread DeMemWM frame memory diffusion metadata

Browse files
algorithms/dememwm/models/diffusion.py CHANGED
@@ -155,8 +155,14 @@ class Diffusion(nn.Module):
155
  def add_shape_channels(self, x):
156
  return rearrange(x, f"... -> ...{' 1' * len(self.x_shape)}")
157
 
 
 
 
 
158
  def model_predictions(self, x, t, action_cond=None, current_frame=None,
159
- pose_cond=None, mode="training", reference_length=None, frame_idx=None):
 
 
160
  x = x.permute(1,0,2,3,4)
161
  action_cond = action_cond.permute(1,0,2)
162
  if pose_cond is not None and pose_cond[0] is not None:
@@ -165,8 +171,21 @@ class Diffusion(nn.Module):
165
  except:
166
  pass
167
  t = t.permute(1,0)
168
- model_output = self.model(x, t, action_cond, current_frame=current_frame, pose_cond=pose_cond,
169
- mode=mode, reference_length=reference_length, frame_idx=frame_idx)
 
 
 
 
 
 
 
 
 
 
 
 
 
170
  model_output = model_output.permute(1,0,2,3,4)
171
  x = x.permute(1,0,2,3,4)
172
  t = t.permute(1,0)
@@ -234,10 +253,14 @@ class Diffusion(nn.Module):
234
  + extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
235
  )
236
 
237
- def p_mean_variance(self, x, t, action_cond=None, pose_cond=None, reference_length=None):
 
 
238
  model_pred = self.model_predictions(x=x, t=t, action_cond=action_cond,
239
  pose_cond=pose_cond, reference_length=reference_length,
240
- frame_idx=frame_idx)
 
 
241
  x_start = model_pred.pred_x_start
242
  return self.q_posterior(x_start=x_start, x_t=x, t=t)
243
 
@@ -286,7 +309,11 @@ class Diffusion(nn.Module):
286
  pose_cond,
287
  noise_levels: torch.Tensor,
288
  reference_length,
289
- frame_idx=None
 
 
 
 
290
  ):
291
  noise = torch.randn_like(x)
292
  noise = torch.clamp(noise, -self.clip_noise, self.clip_noise)
@@ -294,7 +321,10 @@ class Diffusion(nn.Module):
294
  noised_x = self.q_sample(x_start=x, t=noise_levels, noise=noise)
295
 
296
  model_pred = self.model_predictions(x=noised_x, t=noise_levels, action_cond=action_cond,
297
- pose_cond=pose_cond,reference_length=reference_length, frame_idx=frame_idx)
 
 
 
298
 
299
  pred = model_pred.model_out
300
  x_pred = model_pred.pred_x_start
@@ -329,7 +359,11 @@ class Diffusion(nn.Module):
329
  current_frame=None,
330
  mode="training",
331
  reference_length=None,
332
- frame_idx=None
 
 
 
 
333
  ):
334
  real_steps = torch.linspace(-1, self.timesteps - 1, steps=self.sampling_timesteps + 1, device=x.device).long()
335
 
@@ -348,7 +382,11 @@ class Diffusion(nn.Module):
348
  current_frame=current_frame,
349
  mode=mode,
350
  reference_length=reference_length,
351
- frame_idx=frame_idx
 
 
 
 
352
  )
353
 
354
  # FIXME: temporary code for checking ddpm sampling
@@ -367,7 +405,11 @@ class Diffusion(nn.Module):
367
  curr_noise_level=curr_noise_level,
368
  guidance_fn=guidance_fn,
369
  reference_length=reference_length,
370
- frame_idx=frame_idx
 
 
 
 
371
  )
372
 
373
  def ddpm_sample_step(
@@ -379,6 +421,10 @@ class Diffusion(nn.Module):
379
  guidance_fn: Optional[Callable] = None,
380
  reference_length=None,
381
  frame_idx=None,
 
 
 
 
382
  ):
383
  clipped_curr_noise_level = torch.where(
384
  curr_noise_level < 0,
@@ -405,7 +451,11 @@ class Diffusion(nn.Module):
405
  action_cond=action_cond,
406
  pose_cond=pose_cond,
407
  reference_length=reference_length,
408
- frame_idx=frame_idx
 
 
 
 
409
  )
410
 
411
  noise = torch.where(
@@ -430,7 +480,11 @@ class Diffusion(nn.Module):
430
  current_frame=None,
431
  mode="training",
432
  reference_length=None,
433
- frame_idx=None
 
 
 
 
434
  ):
435
  # convert noise level -1 to self.stabilization_level - 1
436
  clipped_curr_noise_level = torch.where(
@@ -477,7 +531,11 @@ class Diffusion(nn.Module):
477
  current_frame=current_frame,
478
  mode=mode,
479
  reference_length=reference_length,
480
- frame_idx=frame_idx
 
 
 
 
481
  )
482
 
483
  guidance_loss = guidance_fn(model_pred.pred_x_start)
@@ -499,7 +557,11 @@ class Diffusion(nn.Module):
499
  current_frame=current_frame,
500
  mode=mode,
501
  reference_length=reference_length,
502
- frame_idx=frame_idx
 
 
 
 
503
  )
504
  x_start = model_pred.pred_x_start
505
  pred_noise = model_pred.pred_noise
 
155
  def add_shape_channels(self, x):
156
  return rearrange(x, f"... -> ...{' 1' * len(self.x_shape)}")
157
 
158
+ @staticmethod
159
+ def _target_frames(x, frame_memory_segments):
160
+ return frame_memory_segments["target"] if frame_memory_segments else x.shape[0]
161
+
162
  def model_predictions(self, x, t, action_cond=None, current_frame=None,
163
+ pose_cond=None, mode="training", reference_length=None, frame_idx=None,
164
+ frame_memory_segments=None, frame_memory_masks=None, frame_memory_pose=None,
165
+ image_hw=None):
166
  x = x.permute(1,0,2,3,4)
167
  action_cond = action_cond.permute(1,0,2)
168
  if pose_cond is not None and pose_cond[0] is not None:
 
171
  except:
172
  pass
173
  t = t.permute(1,0)
174
+ model_kwargs = {
175
+ "current_frame": current_frame,
176
+ "pose_cond": pose_cond,
177
+ "mode": mode,
178
+ "reference_length": reference_length,
179
+ "frame_idx": frame_idx,
180
+ }
181
+ if frame_memory_segments is not None:
182
+ model_kwargs.update({
183
+ "frame_memory_segments": frame_memory_segments,
184
+ "frame_memory_masks": frame_memory_masks,
185
+ "frame_memory_pose": frame_memory_pose,
186
+ "image_hw": image_hw,
187
+ })
188
+ model_output = self.model(x, t, action_cond, **model_kwargs)
189
  model_output = model_output.permute(1,0,2,3,4)
190
  x = x.permute(1,0,2,3,4)
191
  t = t.permute(1,0)
 
253
  + extract(self.sqrt_one_minus_alphas_cumprod, t, x_start.shape) * noise
254
  )
255
 
256
+ def p_mean_variance(self, x, t, action_cond=None, pose_cond=None, reference_length=None,
257
+ frame_idx=None, frame_memory_segments=None, frame_memory_masks=None,
258
+ frame_memory_pose=None, image_hw=None):
259
  model_pred = self.model_predictions(x=x, t=t, action_cond=action_cond,
260
  pose_cond=pose_cond, reference_length=reference_length,
261
+ frame_idx=frame_idx, frame_memory_segments=frame_memory_segments,
262
+ frame_memory_masks=frame_memory_masks,
263
+ frame_memory_pose=frame_memory_pose, image_hw=image_hw)
264
  x_start = model_pred.pred_x_start
265
  return self.q_posterior(x_start=x_start, x_t=x, t=t)
266
 
 
309
  pose_cond,
310
  noise_levels: torch.Tensor,
311
  reference_length,
312
+ frame_idx=None,
313
+ frame_memory_segments=None,
314
+ frame_memory_masks=None,
315
+ frame_memory_pose=None,
316
+ image_hw=None
317
  ):
318
  noise = torch.randn_like(x)
319
  noise = torch.clamp(noise, -self.clip_noise, self.clip_noise)
 
321
  noised_x = self.q_sample(x_start=x, t=noise_levels, noise=noise)
322
 
323
  model_pred = self.model_predictions(x=noised_x, t=noise_levels, action_cond=action_cond,
324
+ pose_cond=pose_cond,reference_length=reference_length, frame_idx=frame_idx,
325
+ frame_memory_segments=frame_memory_segments,
326
+ frame_memory_masks=frame_memory_masks,
327
+ frame_memory_pose=frame_memory_pose, image_hw=image_hw)
328
 
329
  pred = model_pred.model_out
330
  x_pred = model_pred.pred_x_start
 
359
  current_frame=None,
360
  mode="training",
361
  reference_length=None,
362
+ frame_idx=None,
363
+ frame_memory_segments=None,
364
+ frame_memory_masks=None,
365
+ frame_memory_pose=None,
366
+ image_hw=None
367
  ):
368
  real_steps = torch.linspace(-1, self.timesteps - 1, steps=self.sampling_timesteps + 1, device=x.device).long()
369
 
 
382
  current_frame=current_frame,
383
  mode=mode,
384
  reference_length=reference_length,
385
+ frame_idx=frame_idx,
386
+ frame_memory_segments=frame_memory_segments,
387
+ frame_memory_masks=frame_memory_masks,
388
+ frame_memory_pose=frame_memory_pose,
389
+ image_hw=image_hw
390
  )
391
 
392
  # FIXME: temporary code for checking ddpm sampling
 
405
  curr_noise_level=curr_noise_level,
406
  guidance_fn=guidance_fn,
407
  reference_length=reference_length,
408
+ frame_idx=frame_idx,
409
+ frame_memory_segments=frame_memory_segments,
410
+ frame_memory_masks=frame_memory_masks,
411
+ frame_memory_pose=frame_memory_pose,
412
+ image_hw=image_hw
413
  )
414
 
415
  def ddpm_sample_step(
 
421
  guidance_fn: Optional[Callable] = None,
422
  reference_length=None,
423
  frame_idx=None,
424
+ frame_memory_segments=None,
425
+ frame_memory_masks=None,
426
+ frame_memory_pose=None,
427
+ image_hw=None,
428
  ):
429
  clipped_curr_noise_level = torch.where(
430
  curr_noise_level < 0,
 
451
  action_cond=action_cond,
452
  pose_cond=pose_cond,
453
  reference_length=reference_length,
454
+ frame_idx=frame_idx,
455
+ frame_memory_segments=frame_memory_segments,
456
+ frame_memory_masks=frame_memory_masks,
457
+ frame_memory_pose=frame_memory_pose,
458
+ image_hw=image_hw
459
  )
460
 
461
  noise = torch.where(
 
480
  current_frame=None,
481
  mode="training",
482
  reference_length=None,
483
+ frame_idx=None,
484
+ frame_memory_segments=None,
485
+ frame_memory_masks=None,
486
+ frame_memory_pose=None,
487
+ image_hw=None
488
  ):
489
  # convert noise level -1 to self.stabilization_level - 1
490
  clipped_curr_noise_level = torch.where(
 
531
  current_frame=current_frame,
532
  mode=mode,
533
  reference_length=reference_length,
534
+ frame_idx=frame_idx,
535
+ frame_memory_segments=frame_memory_segments,
536
+ frame_memory_masks=frame_memory_masks,
537
+ frame_memory_pose=frame_memory_pose,
538
+ image_hw=image_hw
539
  )
540
 
541
  guidance_loss = guidance_fn(model_pred.pred_x_start)
 
557
  current_frame=current_frame,
558
  mode=mode,
559
  reference_length=reference_length,
560
+ frame_idx=frame_idx,
561
+ frame_memory_segments=frame_memory_segments,
562
+ frame_memory_masks=frame_memory_masks,
563
+ frame_memory_pose=frame_memory_pose,
564
+ image_hw=image_hw
565
  )
566
  x_start = model_pred.pred_x_start
567
  pred_noise = model_pred.pred_noise
algorithms/dememwm/models/dit.py CHANGED
@@ -480,7 +480,8 @@ class DiT(nn.Module):
480
  return imgs
481
 
482
  def forward(self, x, t, action_cond=None, pose_cond=None, current_frame=None, mode=None,
483
- reference_length=None, frame_idx=None):
 
484
  """
485
  Forward pass of DiT.
486
  x: (B, T, C, H, W) tensor of spatial inputs (images or latent representations of images)
 
480
  return imgs
481
 
482
  def forward(self, x, t, action_cond=None, pose_cond=None, current_frame=None, mode=None,
483
+ reference_length=None, frame_idx=None, frame_memory_segments=None,
484
+ frame_memory_masks=None, frame_memory_pose=None, image_hw=None):
485
  """
486
  Forward pass of DiT.
487
  x: (B, T, C, H, W) tensor of spatial inputs (images or latent representations of images)