Add DeMemWM trainability controls
Browse files
algorithms/dememwm/df_video.py
CHANGED
|
@@ -28,6 +28,9 @@ import glob
|
|
| 28 |
# Utility Functions
|
| 29 |
_DEMEMWM_SEGMENT_KEYS = ("target", "anchor", "dynamic", "revisit")
|
| 30 |
_DEMEMWM_STREAM_KEYS = ("anchor", "dynamic", "revisit")
|
|
|
|
|
|
|
|
|
|
| 31 |
|
| 32 |
|
| 33 |
def _cfg_get(cfg, key: str, default=None):
|
|
@@ -38,6 +41,50 @@ def _cfg_get(cfg, key: str, default=None):
|
|
| 38 |
return getattr(cfg, key, default)
|
| 39 |
|
| 40 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 41 |
def _derive_memory_condition_length(cfg) -> int:
|
| 42 |
memory_cfg = _cfg_get(cfg, "memory_selection")
|
| 43 |
if memory_cfg is None:
|
|
@@ -609,6 +656,36 @@ class DeMemWMMinecraft(DiffusionForcingBase):
|
|
| 609 |
if self.require_pose_prediction:
|
| 610 |
self.pose_prediction_model = PosePredictionNet()
|
| 611 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 612 |
def _generate_noise_levels(self, xs: torch.Tensor, masks = None) -> torch.Tensor:
|
| 613 |
"""
|
| 614 |
Generate noise levels for training.
|
|
|
|
| 28 |
# Utility Functions
|
| 29 |
_DEMEMWM_SEGMENT_KEYS = ("target", "anchor", "dynamic", "revisit")
|
| 30 |
_DEMEMWM_STREAM_KEYS = ("anchor", "dynamic", "revisit")
|
| 31 |
+
_DEMEMWM_REFERENCE_ATTN_MARKERS = (".r_attn_anchor.", ".r_attn_dynamic.", ".r_attn_revisit.")
|
| 32 |
+
_DEMEMWM_GEOMETRY_PROJ_MARKERS = (".query_pose_proj.", ".key_pose_proj.")
|
| 33 |
+
_DEMEMWM_ADALN_MLP_MARKERS = (".r_adaLN_modulation.", ".r_mlp.")
|
| 34 |
|
| 35 |
|
| 36 |
def _cfg_get(cfg, key: str, default=None):
|
|
|
|
| 41 |
return getattr(cfg, key, default)
|
| 42 |
|
| 43 |
|
| 44 |
+
def _trainability_cfg(cfg):
|
| 45 |
+
return _cfg_get(cfg, "trainability", {})
|
| 46 |
+
|
| 47 |
+
|
| 48 |
+
def _dememwm_full_dit_active(trainability, global_step: int) -> bool:
|
| 49 |
+
if not bool(_cfg_get(trainability, "train_full_dit", False)):
|
| 50 |
+
return False
|
| 51 |
+
start_step = _cfg_get(trainability, "full_dit_start_step", 0)
|
| 52 |
+
return start_step is None or int(global_step) >= int(start_step)
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def _is_dememwm_memory_trainable_parameter(name: str, trainability) -> bool:
|
| 56 |
+
is_geometry = any(marker in name for marker in _DEMEMWM_GEOMETRY_PROJ_MARKERS)
|
| 57 |
+
if bool(_cfg_get(trainability, "geometry_projections", True)) and is_geometry:
|
| 58 |
+
return True
|
| 59 |
+
if bool(_cfg_get(trainability, "reference_attention", True)) and not is_geometry:
|
| 60 |
+
if any(marker in name for marker in _DEMEMWM_REFERENCE_ATTN_MARKERS):
|
| 61 |
+
return True
|
| 62 |
+
if bool(_cfg_get(trainability, "adaln_mlp", True)):
|
| 63 |
+
return any(marker in name for marker in _DEMEMWM_ADALN_MLP_MARKERS)
|
| 64 |
+
return False
|
| 65 |
+
|
| 66 |
+
|
| 67 |
+
def _apply_dememwm_trainability(diffusion_model, vae, trainability, global_step: int) -> None:
|
| 68 |
+
full_dit_active = _dememwm_full_dit_active(trainability, global_step)
|
| 69 |
+
for name, param in diffusion_model.named_parameters():
|
| 70 |
+
param.requires_grad_(full_dit_active or _is_dememwm_memory_trainable_parameter(name, trainability))
|
| 71 |
+
|
| 72 |
+
if vae is not None:
|
| 73 |
+
freeze_vae = bool(_cfg_get(trainability, "freeze_vae", True))
|
| 74 |
+
for param in vae.parameters():
|
| 75 |
+
param.requires_grad_(not freeze_vae)
|
| 76 |
+
|
| 77 |
+
|
| 78 |
+
def _dememwm_optimizer_parameters(diffusion_model, vae, trainability):
|
| 79 |
+
include_full_dit = bool(_cfg_get(trainability, "train_full_dit", False))
|
| 80 |
+
for name, param in diffusion_model.named_parameters():
|
| 81 |
+
if include_full_dit or _is_dememwm_memory_trainable_parameter(name, trainability):
|
| 82 |
+
yield param
|
| 83 |
+
|
| 84 |
+
if vae is not None and not bool(_cfg_get(trainability, "freeze_vae", True)):
|
| 85 |
+
yield from vae.parameters()
|
| 86 |
+
|
| 87 |
+
|
| 88 |
def _derive_memory_condition_length(cfg) -> int:
|
| 89 |
memory_cfg = _cfg_get(cfg, "memory_selection")
|
| 90 |
if memory_cfg is None:
|
|
|
|
| 656 |
if self.require_pose_prediction:
|
| 657 |
self.pose_prediction_model = PosePredictionNet()
|
| 658 |
|
| 659 |
+
self._apply_trainability()
|
| 660 |
+
|
| 661 |
+
def _global_step_for_trainability(self) -> int:
|
| 662 |
+
try:
|
| 663 |
+
trainer = self.trainer
|
| 664 |
+
except RuntimeError:
|
| 665 |
+
return 0
|
| 666 |
+
return int(getattr(trainer, "global_step", 0) or 0)
|
| 667 |
+
|
| 668 |
+
def _apply_trainability(self) -> None:
|
| 669 |
+
_apply_dememwm_trainability(
|
| 670 |
+
self.diffusion_model,
|
| 671 |
+
getattr(self, "vae", None),
|
| 672 |
+
_trainability_cfg(self.cfg),
|
| 673 |
+
self._global_step_for_trainability(),
|
| 674 |
+
)
|
| 675 |
+
|
| 676 |
+
def configure_optimizers(self):
|
| 677 |
+
trainability = _trainability_cfg(self.cfg)
|
| 678 |
+
self._apply_trainability()
|
| 679 |
+
params = tuple(_dememwm_optimizer_parameters(self.diffusion_model, getattr(self, "vae", None), trainability))
|
| 680 |
+
if not params:
|
| 681 |
+
raise ValueError("DeMemWM trainability selected no optimizer parameters")
|
| 682 |
+
return torch.optim.AdamW(
|
| 683 |
+
params, lr=self.cfg.lr, weight_decay=self.cfg.weight_decay, betas=self.cfg.optimizer_beta
|
| 684 |
+
)
|
| 685 |
+
|
| 686 |
+
def on_train_batch_start(self, batch, batch_idx, dataloader_idx=0) -> None:
|
| 687 |
+
self._apply_trainability()
|
| 688 |
+
|
| 689 |
def _generate_noise_levels(self, xs: torch.Tensor, masks = None) -> torch.Tensor:
|
| 690 |
"""
|
| 691 |
Generate noise levels for training.
|
configurations/algorithm/dememwm_base.yaml
CHANGED
|
@@ -14,6 +14,13 @@ memory_noise:
|
|
| 14 |
revisit_max_fraction: 0.25
|
| 15 |
snr_shift: {enabled: true, anchor: 0.0, dynamic: 0.0, revisit: 0.0}
|
| 16 |
noise_route: {anchor: all, dynamic: all, revisit: all}
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 17 |
|
| 18 |
n_tokens: ${dataset.n_frames}
|
| 19 |
action_cond_dim: ${dataset.action_cond_dim}
|
|
|
|
| 14 |
revisit_max_fraction: 0.25
|
| 15 |
snr_shift: {enabled: true, anchor: 0.0, dynamic: 0.0, revisit: 0.0}
|
| 16 |
noise_route: {anchor: all, dynamic: all, revisit: all}
|
| 17 |
+
trainability:
|
| 18 |
+
freeze_vae: true
|
| 19 |
+
reference_attention: true
|
| 20 |
+
adaln_mlp: true
|
| 21 |
+
geometry_projections: true
|
| 22 |
+
train_full_dit: false
|
| 23 |
+
full_dit_start_step: null
|
| 24 |
|
| 25 |
n_tokens: ${dataset.n_frames}
|
| 26 |
action_cond_dim: ${dataset.action_cond_dim}
|
tests/test_dememwm_latent_dataset.py
CHANGED
|
@@ -5,6 +5,7 @@ from pathlib import Path
|
|
| 5 |
|
| 6 |
import numpy as np
|
| 7 |
import torch
|
|
|
|
| 8 |
from omegaconf import OmegaConf
|
| 9 |
|
| 10 |
from datasets.video.memory_selection import select_memory_indices
|
|
@@ -327,6 +328,64 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
|
|
| 327 |
routed = _apply_memory_route_masks(masks, query_noise, cfg, Diffusion, mode="training")
|
| 328 |
self.assertEqual({key: routed[key].tolist() for key in stream_lengths}, {"anchor": [[True], [False]], "dynamic": [[False], [True]], "revisit": [[True], [True]]})
|
| 329 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 330 |
def test_validation_step_builds_online_memory_from_committed_latents(self):
|
| 331 |
import algorithms.dememwm.df_video as df_video
|
| 332 |
from algorithms.dememwm.df_video import DeMemWMMinecraft, _preprocess_dememwm_latent_batch
|
|
|
|
| 5 |
|
| 6 |
import numpy as np
|
| 7 |
import torch
|
| 8 |
+
from torch import nn
|
| 9 |
from omegaconf import OmegaConf
|
| 10 |
|
| 11 |
from datasets.video.memory_selection import select_memory_indices
|
|
|
|
| 328 |
routed = _apply_memory_route_masks(masks, query_noise, cfg, Diffusion, mode="training")
|
| 329 |
self.assertEqual({key: routed[key].tolist() for key in stream_lengths}, {"anchor": [[True], [False]], "dynamic": [[False], [True]], "revisit": [[True], [True]]})
|
| 330 |
|
| 331 |
+
def test_trainability_controls_memory_groups_full_dit_ramp_and_vae_freeze(self):
|
| 332 |
+
from algorithms.dememwm.df_video import (
|
| 333 |
+
_apply_dememwm_trainability,
|
| 334 |
+
_dememwm_optimizer_parameters,
|
| 335 |
+
_trainability_cfg,
|
| 336 |
+
)
|
| 337 |
+
|
| 338 |
+
class TinyReferenceAttention(nn.Module):
|
| 339 |
+
def __init__(self):
|
| 340 |
+
super().__init__()
|
| 341 |
+
self.to_q = nn.Linear(2, 2, bias=False)
|
| 342 |
+
self.query_pose_proj = nn.Linear(6, 2, bias=False)
|
| 343 |
+
self.key_pose_proj = nn.Linear(6, 2, bias=False)
|
| 344 |
+
|
| 345 |
+
class TinyBlock(nn.Module):
|
| 346 |
+
def __init__(self):
|
| 347 |
+
super().__init__()
|
| 348 |
+
self.s_mlp = nn.Linear(2, 2, bias=False)
|
| 349 |
+
self.r_adaLN_modulation = nn.Sequential(nn.SiLU(), nn.Linear(2, 2, bias=False))
|
| 350 |
+
self.r_mlp = nn.Linear(2, 2, bias=False)
|
| 351 |
+
self.r_attn_anchor = TinyReferenceAttention()
|
| 352 |
+
|
| 353 |
+
class TinyDiffusion(nn.Module):
|
| 354 |
+
def __init__(self):
|
| 355 |
+
super().__init__()
|
| 356 |
+
self.model = nn.Module()
|
| 357 |
+
self.model.blocks = nn.ModuleList([TinyBlock()])
|
| 358 |
+
self.model.final_layer = nn.Linear(2, 2, bias=False)
|
| 359 |
+
|
| 360 |
+
diffusion = TinyDiffusion()
|
| 361 |
+
vae = nn.Linear(2, 2, bias=False)
|
| 362 |
+
cfg = OmegaConf.create({"trainability": {"freeze_vae": True, "train_full_dit": True, "full_dit_start_step": 10, "geometry_projections": False}})
|
| 363 |
+
trainability = _trainability_cfg(cfg)
|
| 364 |
+
|
| 365 |
+
_apply_dememwm_trainability(diffusion, vae, trainability, global_step=0)
|
| 366 |
+
flags = {name: param.requires_grad for name, param in diffusion.named_parameters()}
|
| 367 |
+
self.assertTrue(flags["model.blocks.0.r_attn_anchor.to_q.weight"])
|
| 368 |
+
self.assertFalse(flags["model.blocks.0.r_attn_anchor.query_pose_proj.weight"])
|
| 369 |
+
self.assertFalse(flags["model.blocks.0.r_attn_anchor.key_pose_proj.weight"])
|
| 370 |
+
self.assertTrue(flags["model.blocks.0.r_adaLN_modulation.1.weight"])
|
| 371 |
+
self.assertTrue(flags["model.blocks.0.r_mlp.weight"])
|
| 372 |
+
self.assertFalse(flags["model.blocks.0.s_mlp.weight"])
|
| 373 |
+
self.assertFalse(flags["model.final_layer.weight"])
|
| 374 |
+
self.assertFalse(any(param.requires_grad for param in vae.parameters()))
|
| 375 |
+
|
| 376 |
+
opt_ids = {id(param) for param in _dememwm_optimizer_parameters(diffusion, vae, trainability)}
|
| 377 |
+
self.assertEqual(opt_ids, {id(param) for param in diffusion.parameters()})
|
| 378 |
+
|
| 379 |
+
_apply_dememwm_trainability(diffusion, vae, trainability, global_step=10)
|
| 380 |
+
self.assertTrue(all(param.requires_grad for param in diffusion.parameters()))
|
| 381 |
+
self.assertFalse(any(param.requires_grad for param in vae.parameters()))
|
| 382 |
+
|
| 383 |
+
memory_only = OmegaConf.create({"trainability": {"freeze_vae": True}})
|
| 384 |
+
_apply_dememwm_trainability(diffusion, vae, _trainability_cfg(memory_only), global_step=0)
|
| 385 |
+
default_flags = {name: param.requires_grad for name, param in diffusion.named_parameters()}
|
| 386 |
+
self.assertTrue(default_flags["model.blocks.0.r_attn_anchor.query_pose_proj.weight"])
|
| 387 |
+
self.assertFalse(default_flags["model.final_layer.weight"])
|
| 388 |
+
|
| 389 |
def test_validation_step_builds_online_memory_from_committed_latents(self):
|
| 390 |
import algorithms.dememwm.df_video as df_video
|
| 391 |
from algorithms.dememwm.df_video import DeMemWMMinecraft, _preprocess_dememwm_latent_batch
|