MCPcopy Create free account
hub / github.com/VisionRush/DeepFakeDefenders / train_multi_class

Method train_multi_class

core/mengine.py:70–115  ·  view source on GitHub ↗
(self, train_loader, epoch_idx, ema_start)

Source from the content-addressed store, hash-verified

68 eta_min=cfg.scheduler.min_lr)
69
70 def train_multi_class(self, train_loader, epoch_idx, ema_start):
71 starttime = datetime.datetime.now()
72 # switch to train mode
73 self.netloc_.train()
74 self.loss_meter_.reset()
75 self.top1_meter_.reset()
76 # train
77 train_loader = tqdm(train_loader, desc='train', ascii=True)
78 for imgs_idx, (imgs_tensor, imgs_label, _, _) in enumerate(train_loader):
79 # set cuda
80 imgs_tensor = imgs_tensor.cuda() # [256, 3, 224, 224]
81 imgs_label = imgs_label.cuda()
82 # clear gradients(zero the parameter gradients)
83 self.optimizer_.zero_grad()
84 # calc forward
85 preds = self.netloc_(imgs_tensor)
86 # calc acc & loss
87 loss = self.criterion_(preds, imgs_label)
88
89 # backpropagation
90 loss.backward()
91 # update parameters
92 self.optimizer_.step()
93
94 # EMA update
95 if ema_start:
96 self.ema_model.update(self.netloc_)
97
98 # accumulate loss & acc
99 acc1 = simple_accuracy(preds, imgs_label)
100 if self.DDP:
101 loss = reduce_tensor(loss, self.world_size)
102 acc1 = reduce_tensor(acc1, self.world_size)
103 self.loss_meter_.update(loss.data.item())
104 self.top1_meter_.update(acc1.item())
105
106 # eval
107 top1 = self.top1_meter_.mean
108 loss = self.loss_meter_.mean
109 endtime = datetime.datetime.now()
110 self.lr_ = self.optimizer_.param_groups[0]['lr']
111 if self.local_rank == 0:
112 print('log: epoch-%d, train_top1 is %f, train_loss is %f, lr is %f, time is %d' % (
113 epoch_idx, top1, loss, self.lr_, (endtime - starttime).seconds))
114 # return
115 return top1, loss, self.lr_
116
117 def val_multi_class(self, val_loader, epoch_idx):
118 np.set_printoptions(suppress=True)

Callers 2

main_train.pyFile · 0.80

Calls 4

simple_accuracyFunction · 0.90
reduce_tensorFunction · 0.85
updateMethod · 0.80
resetMethod · 0.45

Tested by

no test coverage detected