(epochs, to_be_distributed)
| 94 | |
| 95 | |
| 96 | def init_models_optimizers(epochs, to_be_distributed): |
| 97 | model = BiRefNet(bb_pretrained=True) |
| 98 | if args.resume: |
| 99 | if os.path.isfile(args.resume): |
| 100 | logger.info("=> loading checkpoint '{}'".format(args.resume)) |
| 101 | state_dict = torch.load(args.resume, map_location='cpu') |
| 102 | state_dict = check_state_dict(state_dict) |
| 103 | model.load_state_dict(state_dict) |
| 104 | epoch_st = int(args.resume.rstrip('.pth').split('epoch_')[-1]) + 1 |
| 105 | else: |
| 106 | logger.info("=> no checkpoint found at '{}'".format(args.resume)) |
| 107 | if to_be_distributed: |
| 108 | model = model.to(device) |
| 109 | model = DDP(model, device_ids=[device]) |
| 110 | else: |
| 111 | model = model.to(device) |
| 112 | if config.compile: |
| 113 | model = torch.compile(model, mode=['default', 'reduce-overhead', 'max-autotune'][0]) |
| 114 | if config.precisionHigh: |
| 115 | torch.set_float32_matmul_precision('high') |
| 116 | |
| 117 | |
| 118 | # Setting optimizer |
| 119 | if config.optimizer == 'AdamW': |
| 120 | optimizer = optim.AdamW(params=model.parameters(), lr=config.lr, weight_decay=1e-2) |
| 121 | elif config.optimizer == 'Adam': |
| 122 | optimizer = optim.Adam(params=model.parameters(), lr=config.lr, weight_decay=0) |
| 123 | lr_scheduler = torch.optim.lr_scheduler.MultiStepLR( |
| 124 | optimizer, |
| 125 | milestones=[lde if lde > 0 else epochs + lde + 1 for lde in config.lr_decay_epochs], |
| 126 | gamma=config.lr_decay_rate |
| 127 | ) |
| 128 | logger.info("Optimizer details:"); logger.info(optimizer) |
| 129 | logger.info("Scheduler details:"); logger.info(lr_scheduler) |
| 130 | |
| 131 | return model, optimizer, lr_scheduler |
| 132 | |
| 133 | |
| 134 | class Trainer: |
no test coverage detected