INNER CODE UNIT · Python
last_epoch
QizhiPei/FABind · FABind/fabind/main_fabind.py:262
last_epoch = -1
steps_per_epoch = len(train_loader)
total_training_steps = args.total_epochs * len(train_loader)
scheduler_warm_up = torch.optim.lr_scheduler.LinearLR(
optimizer,
start_factor=0.5,
end_factor=1,
total_iters=args.warmup_epochs * len(train_loader),
last_epoch=last_epoch,
)
if args.lr_scheduler == "constant":
scheduler_post = torch.optim.lr_scheduler.ConstantLR(optimizer, factor=1.0, total_iters=(args.total_epochs - args.warmup_epochs)*len(train_loader), last_epoch=last_epoch)
elif args.lr_scheduler == "poly_decay":
scheduler_post = torch.optim.lr_scheduler.LinearLR(optimizer, start_factor=1.0, end_factor=0.0, total_iters=(args.total_epochs - args.warmup_epochs)*len(train_loader), last_epoch=last_epoch)
elif args.lr_scheduler == "exp_decay":
scheduler_post = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.995, last_epoch=last_epoch)
elif args.lr_scheduler == "cosine_decay":
scheduler_post = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=(args.total_epochs - args.warmup_epochs)*len(train_loader), eta_min=1e-5, last_epoch=last_epoch)