This commit is contained in:
Rayane 2026-08-16 12:12:35 +08:00 committed by GitHub
commit 39b3d357fe
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
1 changed files with 9 additions and 1 deletions

View File

@ -21,6 +21,14 @@ def get_lr_scheduler(
optimizer, **kwargs
)
elif name == "step":
# StepLR decays purely on step_size/gamma and has no notion of run
# length, so drop the total_iters the trainer injects.
kwargs.pop('total_iters', None)
if 'step_size' not in kwargs:
raise ValueError(
"lr_scheduler 'step' requires lr_scheduler_params.step_size "
"(number of steps between each lr decay)"
)
return torch.optim.lr_scheduler.StepLR(
optimizer, **kwargs
@ -40,7 +48,7 @@ def get_lr_scheduler(
if 'num_warmup_steps' not in kwargs:
print(f"WARNING: num_warmup_steps not in kwargs. Using default value of 1000")
kwargs['num_warmup_steps'] = 1000
del kwargs['total_iters']
kwargs.pop('total_iters', None)
return get_constant_schedule_with_warmup(optimizer, **kwargs)
else:
# try to use a diffusers scheduler