BonanDing commited on
Commit
17a7880
·
1 Parent(s): 0343787

Add DeMemWM optimizer LR groups and warmup

Browse files
algorithms/dememwm/df_base.py CHANGED
@@ -60,14 +60,17 @@ class DiffusionForcingBase(BasePytorchAlgo):
60
  return optimizer_dynamics
61
 
62
  def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_closure):
63
- # update params
64
  optimizer.step(closure=optimizer_closure)
65
-
66
- # manually warm up lr without a scheduler
67
- if self.trainer.global_step < self.cfg.warmup_steps:
68
- lr_scale = min(1.0, float(self.trainer.global_step + 1) / self.cfg.warmup_steps)
69
- for pg in optimizer.param_groups:
70
- pg["lr"] = lr_scale * self.cfg.lr
 
 
 
71
 
72
  def training_step(self, batch, batch_idx) -> STEP_OUTPUT:
73
  xs, conditions, masks = self._preprocess_batch(batch)
 
60
  return optimizer_dynamics
61
 
62
  def optimizer_step(self, epoch, batch_idx, optimizer, optimizer_closure):
63
+ self._apply_optimizer_lrs(optimizer, self.trainer.global_step)
64
  optimizer.step(closure=optimizer_closure)
65
+ self._apply_optimizer_lrs(optimizer, self.trainer.global_step + 1)
66
+
67
+ def _apply_optimizer_lrs(self, optimizer, step: int) -> None:
68
+ warmup_steps = int(getattr(self.cfg, "warmup_steps", 0) or 0)
69
+ for pg in optimizer.param_groups:
70
+ warmup_start_step = int(pg.get("warmup_start_step", 0) or 0)
71
+ warmup_progress = max(0, int(step) - warmup_start_step + 1)
72
+ lr_scale = 1.0 if warmup_steps <= 0 else min(1.0, float(warmup_progress) / warmup_steps)
73
+ pg["lr"] = float(pg.get("target_lr", self.cfg.lr)) * lr_scale
74
 
75
  def training_step(self, batch, batch_idx) -> STEP_OUTPUT:
76
  xs, conditions, masks = self._preprocess_batch(batch)
algorithms/dememwm/df_video.py CHANGED
@@ -103,6 +103,60 @@ def _dememwm_optimizer_parameters(diffusion_model, vae, trainability):
103
  yield from vae.parameters()
104
 
105
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
106
  def _derive_memory_condition_length(cfg) -> int:
107
  memory_cfg = _cfg_get(cfg, "memory_selection")
108
  if memory_cfg is None:
@@ -491,15 +545,24 @@ class DeMemWMMinecraft(DiffusionForcingBase):
491
  def configure_optimizers(self):
492
  trainability = _trainability_cfg(self.cfg)
493
  self._apply_trainability()
494
- params = tuple(_dememwm_optimizer_parameters(self.diffusion_model, getattr(self, "vae", None), trainability))
495
- if not params:
 
 
 
 
 
 
496
  raise ValueError("DeMemWM trainability selected no optimizer parameters")
497
  return torch.optim.AdamW(
498
- params, lr=self.cfg.lr, weight_decay=self.cfg.weight_decay, betas=self.cfg.optimizer_beta
499
  )
500
 
501
  def on_train_batch_start(self, batch, batch_idx, dataloader_idx=0) -> None:
502
  self._apply_trainability()
 
 
 
503
 
504
  def _generate_noise_levels(self, xs: torch.Tensor, masks = None) -> torch.Tensor:
505
  """
 
103
  yield from vae.parameters()
104
 
105
 
106
+ def _dememwm_group_target_lr(group_name: str, trainability, default_lr, global_step: int) -> float:
107
+ lr_cfg = _cfg_get(trainability, "lr", {})
108
+ if group_name == "memory_modules":
109
+ return float(_cfg_get(lr_cfg, "memory_modules", default_lr))
110
+ if group_name == "base_dit":
111
+ if not _dememwm_full_dit_active(trainability, global_step):
112
+ return 0.0
113
+ return float(_cfg_get(lr_cfg, "base_dit", default_lr))
114
+ return float(default_lr)
115
+
116
+
117
+ def _dememwm_group_warmup_start_step(group_name: str, trainability) -> int:
118
+ if group_name != "base_dit":
119
+ return 0
120
+ start_step = _cfg_get(trainability, "full_dit_start_step", 0)
121
+ return 0 if start_step is None else int(start_step)
122
+
123
+
124
+ def _dememwm_optimizer_parameter_groups(diffusion_model, vae, trainability, default_lr, global_step: int):
125
+ include_full_dit = bool(_cfg_get(trainability, "train_full_dit", False))
126
+ grouped = {"memory_modules": [], "base_dit": []}
127
+ for name, param in diffusion_model.named_parameters():
128
+ if _is_dememwm_memory_trainable_parameter(name, trainability):
129
+ grouped["memory_modules"].append(param)
130
+ elif include_full_dit:
131
+ grouped["base_dit"].append(param)
132
+
133
+ param_groups = []
134
+ for group_name, params in grouped.items():
135
+ if params:
136
+ target_lr = _dememwm_group_target_lr(group_name, trainability, default_lr, global_step)
137
+ param_groups.append({
138
+ "params": params,
139
+ "lr": target_lr,
140
+ "target_lr": target_lr,
141
+ "warmup_start_step": _dememwm_group_warmup_start_step(group_name, trainability),
142
+ "name": group_name,
143
+ })
144
+
145
+ if vae is not None and not bool(_cfg_get(trainability, "freeze_vae", True)):
146
+ target_lr = float(_cfg_get(_cfg_get(trainability, "lr", {}), "vae", default_lr))
147
+ param_groups.append({"params": tuple(vae.parameters()), "lr": target_lr, "target_lr": target_lr, "name": "vae"})
148
+
149
+ return param_groups
150
+
151
+
152
+ def _apply_dememwm_optimizer_group_lrs(optimizer, trainability, default_lr, global_step: int) -> None:
153
+ for param_group in optimizer.param_groups:
154
+ group_name = param_group.get("name")
155
+ if group_name in {"memory_modules", "base_dit"}:
156
+ param_group["target_lr"] = _dememwm_group_target_lr(group_name, trainability, default_lr, global_step)
157
+ param_group["warmup_start_step"] = _dememwm_group_warmup_start_step(group_name, trainability)
158
+
159
+
160
  def _derive_memory_condition_length(cfg) -> int:
161
  memory_cfg = _cfg_get(cfg, "memory_selection")
162
  if memory_cfg is None:
 
545
  def configure_optimizers(self):
546
  trainability = _trainability_cfg(self.cfg)
547
  self._apply_trainability()
548
+ param_groups = _dememwm_optimizer_parameter_groups(
549
+ self.diffusion_model,
550
+ getattr(self, "vae", None),
551
+ trainability,
552
+ self.cfg.lr,
553
+ self._global_step_for_trainability(),
554
+ )
555
+ if not param_groups:
556
  raise ValueError("DeMemWM trainability selected no optimizer parameters")
557
  return torch.optim.AdamW(
558
+ param_groups, weight_decay=self.cfg.weight_decay, betas=self.cfg.optimizer_beta
559
  )
560
 
561
  def on_train_batch_start(self, batch, batch_idx, dataloader_idx=0) -> None:
562
  self._apply_trainability()
563
+ trainability = _trainability_cfg(self.cfg)
564
+ for optimizer in getattr(getattr(self, "trainer", None), "optimizers", []) or []:
565
+ _apply_dememwm_optimizer_group_lrs(optimizer, trainability, self.cfg.lr, self._global_step_for_trainability())
566
 
567
  def _generate_noise_levels(self, xs: torch.Tensor, masks = None) -> torch.Tensor:
568
  """
configurations/algorithm/dememwm_base.yaml CHANGED
@@ -23,6 +23,9 @@ trainability:
23
  geometry_projections: true
24
  train_full_dit: false
25
  full_dit_start_step: null
 
 
 
26
 
27
  n_tokens: ${dataset.n_frames}
28
  action_cond_dim: ${dataset.action_cond_dim}
 
23
  geometry_projections: true
24
  train_full_dit: false
25
  full_dit_start_step: null
26
+ lr:
27
+ memory_modules: 8.0e-5
28
+ base_dit: 2.0e-5
29
 
30
  n_tokens: ${dataset.n_frames}
31
  action_cond_dim: ${dataset.action_cond_dim}
tests/test_dememwm_latent_dataset.py CHANGED
@@ -383,8 +383,8 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
383
  def reweight_loss(self, loss, weight=None):
384
  raise AssertionError("target mask path should reduce the latent dict loss directly")
385
 
386
- def log(self, name, value):
387
- self.logged.append((name, value))
388
 
389
  batch = {
390
  "latents": torch.arange(5, dtype=torch.float32).view(1, 5, 1, 1, 1),
@@ -492,9 +492,12 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
492
  def test_trainability_controls_memory_groups_full_dit_ramp_and_vae_freeze(self):
493
  from algorithms.dememwm.df_video import (
494
  _apply_dememwm_trainability,
 
 
495
  _dememwm_optimizer_parameters,
496
  _trainability_cfg,
497
  )
 
498
 
499
  class TinyReferenceAttention(nn.Module):
500
  def __init__(self):
@@ -520,7 +523,15 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
520
 
521
  diffusion = TinyDiffusion()
522
  vae = nn.Linear(2, 2, bias=False)
523
- cfg = OmegaConf.create({"trainability": {"freeze_vae": True, "train_full_dit": True, "full_dit_start_step": 10, "geometry_projections": False}})
 
 
 
 
 
 
 
 
524
  trainability = _trainability_cfg(cfg)
525
 
526
  _apply_dememwm_trainability(diffusion, vae, trainability, global_step=0)
@@ -536,6 +547,34 @@ class DeMemWMLatentDatasetTests(unittest.TestCase):
536
 
537
  opt_ids = {id(param) for param in _dememwm_optimizer_parameters(diffusion, vae, trainability)}
538
  self.assertEqual(opt_ids, {id(param) for param in diffusion.parameters()})
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
539
 
540
  _apply_dememwm_trainability(diffusion, vae, trainability, global_step=10)
541
  self.assertTrue(all(param.requires_grad for param in diffusion.parameters()))
 
383
  def reweight_loss(self, loss, weight=None):
384
  raise AssertionError("target mask path should reduce the latent dict loss directly")
385
 
386
+ def log(self, name, value, **kwargs):
387
+ self.logged.append((name, value, kwargs))
388
 
389
  batch = {
390
  "latents": torch.arange(5, dtype=torch.float32).view(1, 5, 1, 1, 1),
 
492
  def test_trainability_controls_memory_groups_full_dit_ramp_and_vae_freeze(self):
493
  from algorithms.dememwm.df_video import (
494
  _apply_dememwm_trainability,
495
+ _apply_dememwm_optimizer_group_lrs,
496
+ _dememwm_optimizer_parameter_groups,
497
  _dememwm_optimizer_parameters,
498
  _trainability_cfg,
499
  )
500
+ from algorithms.dememwm.df_base import DiffusionForcingBase
501
 
502
  class TinyReferenceAttention(nn.Module):
503
  def __init__(self):
 
523
 
524
  diffusion = TinyDiffusion()
525
  vae = nn.Linear(2, 2, bias=False)
526
+ cfg = OmegaConf.create({
527
+ "trainability": {
528
+ "freeze_vae": True,
529
+ "train_full_dit": True,
530
+ "full_dit_start_step": 10,
531
+ "geometry_projections": False,
532
+ "lr": {"memory_modules": 8.0e-5, "base_dit": 2.0e-5},
533
+ }
534
+ })
535
  trainability = _trainability_cfg(cfg)
536
 
537
  _apply_dememwm_trainability(diffusion, vae, trainability, global_step=0)
 
547
 
548
  opt_ids = {id(param) for param in _dememwm_optimizer_parameters(diffusion, vae, trainability)}
549
  self.assertEqual(opt_ids, {id(param) for param in diffusion.parameters()})
550
+ param_groups = _dememwm_optimizer_parameter_groups(diffusion, vae, trainability, default_lr=2.0e-5, global_step=0)
551
+ lr_by_name = {group["name"]: group["target_lr"] for group in param_groups}
552
+ warmup_start_by_name = {group["name"]: group.get("warmup_start_step", 0) for group in param_groups}
553
+ self.assertEqual(set(lr_by_name), {"memory_modules", "base_dit"})
554
+ self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-5)
555
+ self.assertEqual(lr_by_name["base_dit"], 0.0)
556
+ self.assertEqual(warmup_start_by_name["memory_modules"], 0)
557
+ self.assertEqual(warmup_start_by_name["base_dit"], 10)
558
+
559
+ optimizer = torch.optim.AdamW(param_groups)
560
+ lr_owner = object.__new__(DiffusionForcingBase)
561
+ lr_owner.cfg = OmegaConf.create({"lr": 2.0e-5, "warmup_steps": 10})
562
+ lr_owner._apply_optimizer_lrs(optimizer, step=0)
563
+ lr_by_name = {group["name"]: group["lr"] for group in optimizer.param_groups}
564
+ self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-6)
565
+ self.assertEqual(lr_by_name["base_dit"], 0.0)
566
+
567
+ _apply_dememwm_optimizer_group_lrs(optimizer, trainability, default_lr=2.0e-5, global_step=10)
568
+ lr_owner._apply_optimizer_lrs(optimizer, step=10)
569
+ lr_by_name = {group["name"]: group["lr"] for group in optimizer.param_groups}
570
+ self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-5)
571
+ self.assertAlmostEqual(lr_by_name["base_dit"], 2.0e-6)
572
+
573
+ _apply_dememwm_optimizer_group_lrs(optimizer, trainability, default_lr=2.0e-5, global_step=19)
574
+ lr_owner._apply_optimizer_lrs(optimizer, step=19)
575
+ lr_by_name = {group["name"]: group["lr"] for group in optimizer.param_groups}
576
+ self.assertAlmostEqual(lr_by_name["memory_modules"], 8.0e-5)
577
+ self.assertAlmostEqual(lr_by_name["base_dit"], 2.0e-5)
578
 
579
  _apply_dememwm_trainability(diffusion, vae, trainability, global_step=10)
580
  self.assertTrue(all(param.requires_grad for param in diffusion.parameters()))