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 |
-
|
| 64 |
optimizer.step(closure=optimizer_closure)
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 495 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 496 |
raise ValueError("DeMemWM trainability selected no optimizer parameters")
|
| 497 |
return torch.optim.AdamW(
|
| 498 |
-
|
| 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({
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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()))
|