BonanDing commited on
Commit
e66c1dc
·
1 Parent(s): 4050b8e

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