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 |
-
|
| 169 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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)
|