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

Class Trainer

train.py:134–287  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

132
133
134class Trainer:
135 def __init__(
136 self, data_loaders, model_opt_lrsch,
137 ):
138 self.model, self.optimizer, self.lr_scheduler = model_opt_lrsch
139 self.train_loader, self.test_loaders = data_loaders
140 if config.out_ref:
141 self.criterion_gdt = nn.BCELoss()
142
143 # Setting Losses
144 self.pix_loss = PixLoss()
145 self.cls_loss = ClsLoss()
146
147 # Others
148 self.loss_log = AverageMeter()
149 if config.lambda_adv_g:
150 self.optimizer_d, self.lr_scheduler_d, self.disc, self.adv_criterion = self._load_adv_components()
151 self.disc_update_for_odd = 0
152
153 def _load_adv_components(self):
154 # AIL
155 from loss import Discriminator
156 disc = Discriminator(channels=3, img_size=config.size)
157 if to_be_distributed:
158 disc = disc.to(device)
159 disc = DDP(disc, device_ids=[device], broadcast_buffers=False)
160 else:
161 disc = disc.to(device)
162 if config.compile:
163 disc = torch.compile(disc, mode=['default', 'reduce-overhead', 'max-autotune'][0])
164 adv_criterion = nn.BCELoss()
165 if config.optimizer == 'AdamW':
166 optimizer_d = optim.AdamW(params=disc.parameters(), lr=config.lr, weight_decay=1e-2)
167 elif config.optimizer == 'Adam':
168 optimizer_d = optim.Adam(params=disc.parameters(), lr=config.lr, weight_decay=0)
169 lr_scheduler_d = torch.optim.lr_scheduler.MultiStepLR(
170 optimizer_d,
171 milestones=[lde if lde > 0 else args.epochs + lde + 1 for lde in config.lr_decay_epochs],
172 gamma=config.lr_decay_rate
173 )
174 return optimizer_d, lr_scheduler_d, disc, adv_criterion
175
176 def _train_batch(self, batch):
177 inputs = batch[0].to(device)
178 gts = batch[1].to(device)
179 class_labels = batch[2].to(device)
180 scaled_preds, class_preds_lst = self.model(inputs)
181 if config.out_ref:
182 (outs_gdt_pred, outs_gdt_label), scaled_preds = scaled_preds
183 for _idx, (_gdt_pred, _gdt_label) in enumerate(zip(outs_gdt_pred, outs_gdt_label)):
184 _gdt_pred = nn.functional.interpolate(_gdt_pred, size=_gdt_label.shape[2:], mode='bilinear', align_corners=True).sigmoid()
185 _gdt_label = _gdt_label.sigmoid()
186 loss_gdt = self.criterion_gdt(_gdt_pred, _gdt_label) if _idx == 0 else self.criterion_gdt(_gdt_pred, _gdt_label) + loss_gdt
187 # self.loss_dict['loss_gdt'] = loss_gdt.item()
188 if None in class_preds_lst:
189 loss_cls = 0.
190 else:
191 loss_cls = self.cls_loss(class_preds_lst, class_labels) * 1.0

Callers 1

mainFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected