MCPcopy Create free account
hub / github.com/ZhengPeng7/BiRefNet / init_models_optimizers

Function init_models_optimizers

train.py:96–131  ·  view source on GitHub ↗
(epochs, to_be_distributed)

Source from the content-addressed store, hash-verified

94
95
96def 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
134class Trainer:

Callers 1

mainFunction · 0.85

Calls 3

BiRefNetClass · 0.90
check_state_dictFunction · 0.90
infoMethod · 0.80

Tested by

no test coverage detected