From 589575985f36ad047b8617ba381ce9f96c0a0dbf Mon Sep 17 00:00:00 2001 From: Musisoul Date: Tue, 4 Aug 2026 12:06:00 +0000 Subject: [PATCH 1/2] fix(fastwam): make LR schedule configurable --- .../libero_plus_uncond_2cam224_4gpu.yaml | 2 + .../libero_plus_uncond_2cam224_8gpu.yaml | 2 + .../train/fastwam/libero_uncond_2cam224.yaml | 2 + .../lightx2v_train/trainers/fastwam.py | 43 ++++-- .../tests/test_fastwam_lr_scheduler.py | 125 ++++++++++++++++++ 5 files changed, 163 insertions(+), 11 deletions(-) create mode 100644 lightx2v_train/tests/test_fastwam_lr_scheduler.py diff --git a/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_4gpu.yaml b/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_4gpu.yaml index 75197e9eb..efc546dcc 100644 --- a/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_4gpu.yaml +++ b/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_4gpu.yaml @@ -31,6 +31,8 @@ training: save_final: true lr_scheduler: cosine lr_warmup_iters: 8742 + # Cosine floor = optimizer.learning_rate * lr_eta_min_ratio; set to 0 to decay to zero. + lr_eta_min_ratio: 0.01 optimizer: learning_rate: 1.0e-4 weight_decay: 1.0e-2 diff --git a/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_8gpu.yaml b/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_8gpu.yaml index 49287b4af..eb6388d8d 100644 --- a/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_8gpu.yaml +++ b/lightx2v_train/configs/train/fastwam/libero_plus_uncond_2cam224_8gpu.yaml @@ -31,6 +31,8 @@ training: save_final: true lr_scheduler: cosine lr_warmup_iters: 8742 + # Cosine floor = optimizer.learning_rate * lr_eta_min_ratio; set to 0 to decay to zero. + lr_eta_min_ratio: 0.01 optimizer: learning_rate: 1.0e-4 weight_decay: 1.0e-2 diff --git a/lightx2v_train/configs/train/fastwam/libero_uncond_2cam224.yaml b/lightx2v_train/configs/train/fastwam/libero_uncond_2cam224.yaml index 44c09802d..4f20d265e 100644 --- a/lightx2v_train/configs/train/fastwam/libero_uncond_2cam224.yaml +++ b/lightx2v_train/configs/train/fastwam/libero_uncond_2cam224.yaml @@ -29,6 +29,8 @@ training: save_final: true lr_scheduler: cosine lr_warmup_iters: 1085 + # Cosine floor = optimizer.learning_rate * lr_eta_min_ratio; set to 0 to decay to zero. + lr_eta_min_ratio: 0.01 optimizer: learning_rate: 1.0e-4 weight_decay: 1.0e-2 diff --git a/lightx2v_train/lightx2v_train/trainers/fastwam.py b/lightx2v_train/lightx2v_train/trainers/fastwam.py index f39d7cc08..3317e0b1e 100644 --- a/lightx2v_train/lightx2v_train/trainers/fastwam.py +++ b/lightx2v_train/lightx2v_train/trainers/fastwam.py @@ -9,7 +9,7 @@ from PIL import Image, ImageDraw from loguru import logger from torch.nn.parallel import DistributedDataParallel -from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR, SequentialLR +from torch.optim.lr_scheduler import ConstantLR, CosineAnnealingLR, LinearLR, SequentialLR from lightx2v_train.runtime.checkpoint import find_latest_checkpoint, parse_checkpoint_iteration, prune_checkpoints from lightx2v_train.runtime.distributed import ( @@ -85,8 +85,9 @@ def __init__(self, config): self.save_every_iters = int(self.training_config.get("save_every_iters", 0) or 0) self.save_total_limit = int(self.training_config.get("save_total_limit", 3)) self.save_final = bool(self.training_config.get("save_final", True)) - self.lr_scheduler_name = self.training_config.get("lr_scheduler", "constant") + self.lr_scheduler_name = str(self.training_config.get("lr_scheduler", "constant")).strip().lower() self.lr_warmup_iters = int(self.training_config.get("lr_warmup_iters", 0)) + self.lr_eta_min_ratio = float(self.training_config.get("lr_eta_min_ratio", 0.01)) self.train_log_every_iters = max(1, int(self.logging_config.get("train_log_every_iters", 10))) zero1_config = self.config.get("distributed", {}).get("zero1", {}) @@ -153,23 +154,43 @@ def _build_optimizer(self): return torch.optim.AdamW(self.trainable_params, **optimizer_kwargs) def _build_lr_scheduler(self): + if self.lr_scheduler_name not in {"constant", "cosine"}: + raise ValueError(f"Unsupported FastWAM lr_scheduler {self.lr_scheduler_name!r}; expected 'constant' or 'cosine'.") + if not 0 < self.lr_warmup_iters < self.max_train_iters: + raise ValueError( + f"FastWAM requires 0 < lr_warmup_iters < max_train_iters for SequentialLR, got " + f"lr_warmup_iters={self.lr_warmup_iters} and max_train_iters={self.max_train_iters}." + ) + if not 0.0 <= self.lr_eta_min_ratio <= 1.0: + raise ValueError(f"FastWAM lr_eta_min_ratio must be in [0, 1], got {self.lr_eta_min_ratio}.") + + # LinearLR changes the optimizer's current LR during construction. Capture + # the configured base LR first so cosine eta_min is not scaled by warmup. + base_lr = float(self.optimizer.param_groups[0]["lr"]) warmup_scheduler = LinearLR( self.optimizer, - start_factor=1e-8, + start_factor=1.0 / self.lr_warmup_iters, end_factor=1.0, total_iters=self.lr_warmup_iters, ) - cosine_scheduler = CosineAnnealingLR( - self.optimizer, - T_max=self.max_train_iters - self.lr_warmup_iters, - eta_min=self.optimizer.param_groups[0]["lr"] * 0.01, - ) - scheduler = SequentialLR( + remaining_iters = self.max_train_iters - self.lr_warmup_iters + if self.lr_scheduler_name == "cosine": + main_scheduler = CosineAnnealingLR( + self.optimizer, + T_max=remaining_iters, + eta_min=base_lr * self.lr_eta_min_ratio, + ) + else: + main_scheduler = ConstantLR( + self.optimizer, + factor=1.0, + total_iters=remaining_iters, + ) + return SequentialLR( self.optimizer, - schedulers=[warmup_scheduler, cosine_scheduler], + schedulers=[warmup_scheduler, main_scheduler], milestones=[self.lr_warmup_iters], ) - return scheduler def setup(self, resume_ckpt_path=None): self.model.set_dit_only_trainable() diff --git a/lightx2v_train/tests/test_fastwam_lr_scheduler.py b/lightx2v_train/tests/test_fastwam_lr_scheduler.py new file mode 100644 index 000000000..5c5992dfd --- /dev/null +++ b/lightx2v_train/tests/test_fastwam_lr_scheduler.py @@ -0,0 +1,125 @@ +import io + +import pytest +import torch +from lightx2v_train.trainers.fastwam import FastWAMTrainer +from torch.optim.lr_scheduler import ConstantLR, CosineAnnealingLR, SequentialLR + + +def _build_trainer(*, scheduler_name="cosine", eta_min_ratio=0.01, max_train_iters=20, warmup_iters=4): + trainer = FastWAMTrainer( + { + "model": {}, + "training": { + "output_dir": "unused", + "max_train_iters": max_train_iters, + "lr_scheduler": scheduler_name, + "lr_warmup_iters": warmup_iters, + "lr_eta_min_ratio": eta_min_ratio, + "optimizer": {"learning_rate": 1.0e-4}, + }, + } + ) + parameter = torch.nn.Parameter(torch.tensor(1.0)) + trainer.optimizer = torch.optim.AdamW([parameter], lr=trainer.learning_rate) + return trainer, parameter + + +def _step(optimizer, scheduler, count): + for _ in range(count): + optimizer.step() + scheduler.step() + + +def test_cosine_scheduler_uses_configured_base_lr_for_eta_min(): + trainer, _ = _build_trainer() + + scheduler = trainer._build_lr_scheduler() + + assert isinstance(scheduler, SequentialLR) + assert isinstance(scheduler._schedulers[1], CosineAnnealingLR) + assert scheduler._schedulers[1].eta_min == pytest.approx(1.0e-6) + assert scheduler.get_last_lr() == pytest.approx([trainer.learning_rate / trainer.lr_warmup_iters]) + + _step(trainer.optimizer, scheduler, trainer.max_train_iters) + assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(1.0e-6) + + +def test_cosine_scheduler_allows_zero_eta_min_ratio(): + trainer, _ = _build_trainer(eta_min_ratio=0.0) + scheduler = trainer._build_lr_scheduler() + + assert scheduler._schedulers[1].eta_min == 0.0 + + _step(trainer.optimizer, scheduler, trainer.max_train_iters) + assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(0.0, abs=1.0e-15) + + +def test_constant_scheduler_keeps_base_lr_after_warmup(): + trainer, _ = _build_trainer(scheduler_name="constant") + scheduler = trainer._build_lr_scheduler() + + assert isinstance(scheduler, SequentialLR) + assert isinstance(scheduler._schedulers[1], ConstantLR) + + _step(trainer.optimizer, scheduler, trainer.lr_warmup_iters) + assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(1.0e-4) + + _step(trainer.optimizer, scheduler, trainer.max_train_iters - trainer.lr_warmup_iters) + assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(1.0e-4) + + +@pytest.mark.parametrize("scheduler_name", ["cosine", "constant"]) +def test_new_sequential_lr_checkpoint_round_trip_continues_same_curve(scheduler_name): + trainer, _ = _build_trainer(scheduler_name=scheduler_name, max_train_iters=12, warmup_iters=3) + scheduler = trainer._build_lr_scheduler() + _step(trainer.optimizer, scheduler, 7) + + checkpoint_buffer = io.BytesIO() + torch.save( + { + "optimizer": trainer.optimizer.state_dict(), + "lr_scheduler": scheduler.state_dict(), + }, + checkpoint_buffer, + ) + checkpoint_buffer.seek(0) + checkpoint = torch.load(checkpoint_buffer, map_location="cpu", weights_only=False) + + resumed_trainer, _ = _build_trainer(scheduler_name=scheduler_name, max_train_iters=12, warmup_iters=3) + resumed_scheduler = resumed_trainer._build_lr_scheduler() + resumed_trainer.optimizer.load_state_dict(checkpoint["optimizer"]) + resumed_scheduler.load_state_dict(checkpoint["lr_scheduler"]) + + assert resumed_scheduler.last_epoch == scheduler.last_epoch + assert resumed_scheduler.get_last_lr() == pytest.approx(scheduler.get_last_lr()) + if scheduler_name == "cosine": + assert resumed_scheduler._schedulers[1].eta_min == pytest.approx(1.0e-6) + + for _ in range(5): + trainer.optimizer.step() + scheduler.step() + resumed_trainer.optimizer.step() + resumed_scheduler.step() + assert resumed_scheduler.get_last_lr() == pytest.approx(scheduler.get_last_lr()) + + +@pytest.mark.parametrize( + ("scheduler_name", "eta_min_ratio", "warmup_iters"), + [ + ("linear", 0.01, 4), + ("cosine", -0.01, 4), + ("cosine", 1.01, 4), + ("cosine", 0.01, 0), + ("cosine", 0.01, 20), + ], +) +def test_invalid_scheduler_configuration_fails_early(scheduler_name, eta_min_ratio, warmup_iters): + trainer, _ = _build_trainer( + scheduler_name=scheduler_name, + eta_min_ratio=eta_min_ratio, + warmup_iters=warmup_iters, + ) + + with pytest.raises(ValueError): + trainer._build_lr_scheduler() From d85ef403c73c860a1d5eeb73947b2c20a07cb334 Mon Sep 17 00:00:00 2001 From: Musisoul Date: Tue, 4 Aug 2026 12:38:16 +0000 Subject: [PATCH 2/2] u --- .../lightx2v_train/trainers/fastwam.py | 28 ++-- .../tests/test_fastwam_lr_scheduler.py | 125 ------------------ 2 files changed, 11 insertions(+), 142 deletions(-) delete mode 100644 lightx2v_train/tests/test_fastwam_lr_scheduler.py diff --git a/lightx2v_train/lightx2v_train/trainers/fastwam.py b/lightx2v_train/lightx2v_train/trainers/fastwam.py index 3317e0b1e..2ac5a84a3 100644 --- a/lightx2v_train/lightx2v_train/trainers/fastwam.py +++ b/lightx2v_train/lightx2v_train/trainers/fastwam.py @@ -85,7 +85,7 @@ def __init__(self, config): self.save_every_iters = int(self.training_config.get("save_every_iters", 0) or 0) self.save_total_limit = int(self.training_config.get("save_total_limit", 3)) self.save_final = bool(self.training_config.get("save_final", True)) - self.lr_scheduler_name = str(self.training_config.get("lr_scheduler", "constant")).strip().lower() + self.lr_scheduler_name = self.training_config.get("lr_scheduler", "constant") self.lr_warmup_iters = int(self.training_config.get("lr_warmup_iters", 0)) self.lr_eta_min_ratio = float(self.training_config.get("lr_eta_min_ratio", 0.01)) self.train_log_every_iters = max(1, int(self.logging_config.get("train_log_every_iters", 10))) @@ -154,25 +154,9 @@ def _build_optimizer(self): return torch.optim.AdamW(self.trainable_params, **optimizer_kwargs) def _build_lr_scheduler(self): - if self.lr_scheduler_name not in {"constant", "cosine"}: - raise ValueError(f"Unsupported FastWAM lr_scheduler {self.lr_scheduler_name!r}; expected 'constant' or 'cosine'.") - if not 0 < self.lr_warmup_iters < self.max_train_iters: - raise ValueError( - f"FastWAM requires 0 < lr_warmup_iters < max_train_iters for SequentialLR, got " - f"lr_warmup_iters={self.lr_warmup_iters} and max_train_iters={self.max_train_iters}." - ) - if not 0.0 <= self.lr_eta_min_ratio <= 1.0: - raise ValueError(f"FastWAM lr_eta_min_ratio must be in [0, 1], got {self.lr_eta_min_ratio}.") - # LinearLR changes the optimizer's current LR during construction. Capture # the configured base LR first so cosine eta_min is not scaled by warmup. base_lr = float(self.optimizer.param_groups[0]["lr"]) - warmup_scheduler = LinearLR( - self.optimizer, - start_factor=1.0 / self.lr_warmup_iters, - end_factor=1.0, - total_iters=self.lr_warmup_iters, - ) remaining_iters = self.max_train_iters - self.lr_warmup_iters if self.lr_scheduler_name == "cosine": main_scheduler = CosineAnnealingLR( @@ -186,6 +170,16 @@ def _build_lr_scheduler(self): factor=1.0, total_iters=remaining_iters, ) + + if self.lr_warmup_iters == 0: + return main_scheduler + + warmup_scheduler = LinearLR( + self.optimizer, + start_factor=1.0 / self.lr_warmup_iters, + end_factor=1.0, + total_iters=self.lr_warmup_iters, + ) return SequentialLR( self.optimizer, schedulers=[warmup_scheduler, main_scheduler], diff --git a/lightx2v_train/tests/test_fastwam_lr_scheduler.py b/lightx2v_train/tests/test_fastwam_lr_scheduler.py deleted file mode 100644 index 5c5992dfd..000000000 --- a/lightx2v_train/tests/test_fastwam_lr_scheduler.py +++ /dev/null @@ -1,125 +0,0 @@ -import io - -import pytest -import torch -from lightx2v_train.trainers.fastwam import FastWAMTrainer -from torch.optim.lr_scheduler import ConstantLR, CosineAnnealingLR, SequentialLR - - -def _build_trainer(*, scheduler_name="cosine", eta_min_ratio=0.01, max_train_iters=20, warmup_iters=4): - trainer = FastWAMTrainer( - { - "model": {}, - "training": { - "output_dir": "unused", - "max_train_iters": max_train_iters, - "lr_scheduler": scheduler_name, - "lr_warmup_iters": warmup_iters, - "lr_eta_min_ratio": eta_min_ratio, - "optimizer": {"learning_rate": 1.0e-4}, - }, - } - ) - parameter = torch.nn.Parameter(torch.tensor(1.0)) - trainer.optimizer = torch.optim.AdamW([parameter], lr=trainer.learning_rate) - return trainer, parameter - - -def _step(optimizer, scheduler, count): - for _ in range(count): - optimizer.step() - scheduler.step() - - -def test_cosine_scheduler_uses_configured_base_lr_for_eta_min(): - trainer, _ = _build_trainer() - - scheduler = trainer._build_lr_scheduler() - - assert isinstance(scheduler, SequentialLR) - assert isinstance(scheduler._schedulers[1], CosineAnnealingLR) - assert scheduler._schedulers[1].eta_min == pytest.approx(1.0e-6) - assert scheduler.get_last_lr() == pytest.approx([trainer.learning_rate / trainer.lr_warmup_iters]) - - _step(trainer.optimizer, scheduler, trainer.max_train_iters) - assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(1.0e-6) - - -def test_cosine_scheduler_allows_zero_eta_min_ratio(): - trainer, _ = _build_trainer(eta_min_ratio=0.0) - scheduler = trainer._build_lr_scheduler() - - assert scheduler._schedulers[1].eta_min == 0.0 - - _step(trainer.optimizer, scheduler, trainer.max_train_iters) - assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(0.0, abs=1.0e-15) - - -def test_constant_scheduler_keeps_base_lr_after_warmup(): - trainer, _ = _build_trainer(scheduler_name="constant") - scheduler = trainer._build_lr_scheduler() - - assert isinstance(scheduler, SequentialLR) - assert isinstance(scheduler._schedulers[1], ConstantLR) - - _step(trainer.optimizer, scheduler, trainer.lr_warmup_iters) - assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(1.0e-4) - - _step(trainer.optimizer, scheduler, trainer.max_train_iters - trainer.lr_warmup_iters) - assert trainer.optimizer.param_groups[0]["lr"] == pytest.approx(1.0e-4) - - -@pytest.mark.parametrize("scheduler_name", ["cosine", "constant"]) -def test_new_sequential_lr_checkpoint_round_trip_continues_same_curve(scheduler_name): - trainer, _ = _build_trainer(scheduler_name=scheduler_name, max_train_iters=12, warmup_iters=3) - scheduler = trainer._build_lr_scheduler() - _step(trainer.optimizer, scheduler, 7) - - checkpoint_buffer = io.BytesIO() - torch.save( - { - "optimizer": trainer.optimizer.state_dict(), - "lr_scheduler": scheduler.state_dict(), - }, - checkpoint_buffer, - ) - checkpoint_buffer.seek(0) - checkpoint = torch.load(checkpoint_buffer, map_location="cpu", weights_only=False) - - resumed_trainer, _ = _build_trainer(scheduler_name=scheduler_name, max_train_iters=12, warmup_iters=3) - resumed_scheduler = resumed_trainer._build_lr_scheduler() - resumed_trainer.optimizer.load_state_dict(checkpoint["optimizer"]) - resumed_scheduler.load_state_dict(checkpoint["lr_scheduler"]) - - assert resumed_scheduler.last_epoch == scheduler.last_epoch - assert resumed_scheduler.get_last_lr() == pytest.approx(scheduler.get_last_lr()) - if scheduler_name == "cosine": - assert resumed_scheduler._schedulers[1].eta_min == pytest.approx(1.0e-6) - - for _ in range(5): - trainer.optimizer.step() - scheduler.step() - resumed_trainer.optimizer.step() - resumed_scheduler.step() - assert resumed_scheduler.get_last_lr() == pytest.approx(scheduler.get_last_lr()) - - -@pytest.mark.parametrize( - ("scheduler_name", "eta_min_ratio", "warmup_iters"), - [ - ("linear", 0.01, 4), - ("cosine", -0.01, 4), - ("cosine", 1.01, 4), - ("cosine", 0.01, 0), - ("cosine", 0.01, 20), - ], -) -def test_invalid_scheduler_configuration_fails_early(scheduler_name, eta_min_ratio, warmup_iters): - trainer, _ = _build_trainer( - scheduler_name=scheduler_name, - eta_min_ratio=eta_min_ratio, - warmup_iters=warmup_iters, - ) - - with pytest.raises(ValueError): - trainer._build_lr_scheduler()