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..2ac5a84a3 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 ( @@ -87,6 +87,7 @@ def __init__(self, config): self.save_final = bool(self.training_config.get("save_final", True)) 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))) zero1_config = self.config.get("distributed", {}).get("zero1", {}) @@ -153,23 +154,37 @@ def _build_optimizer(self): return torch.optim.AdamW(self.trainable_params, **optimizer_kwargs) def _build_lr_scheduler(self): + # 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"]) + 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, + ) + + if self.lr_warmup_iters == 0: + return main_scheduler + 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( + 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()