INNER CODE UNIT · Python

total_training_steps

QizhiPei/FABind · FABind/fabind/main_fabind.py:264

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)
elif args.lr_scheduler == "cosine_decay_restart":
    scheduler_post = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=10, eta_min=0.0001, last_epoch=last_epoch)

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…