MCPcopy Create free account
hub / github.com/XLearning-SCU/2022-CVPR-DART / eval_train

Function eval_train

run.py:290–336  ·  view source on GitHub ↗
(net, dataloader, type)

Source from the content-addressed store, hash-verified

288
289
290def eval_train(net, dataloader, type):
291 losses_V_aug1 = -1. * torch.ones(len(evaltrainset.train_color_label))
292 losses_V_aug2 = -1. * torch.ones(len(evaltrainset.train_color_label))
293 losses_I = -1. * torch.ones(len(evaltrainset.train_thermal_label))
294
295 net.train()
296 with torch.no_grad():
297 for batch_idx, (input10, input11, input2, label1, label2, index_V, index_I) in enumerate(dataloader):
298 input1 = torch.cat((input10, input11,), 0)
299 input1 = input1.cuda()
300 input2 = input2.cuda()
301 label1 = label1.cuda()
302 label2 = label2.cuda()
303
304 index_V = np.concatenate((index_V, index_V), 0)
305 labels = torch.cat((label1, label1, label2), 0)
306 _, out0, = net(input1, input2)
307 loss = criterion_CE(out0, labels)
308 loss1 = loss[0:64]
309 loss2 = loss[64:96]
310
311 for n1 in range(input2.size(0)):
312 losses_V_aug1[index_V[n1]] = loss1[n1]
313 losses_V_aug2[index_V[n1 + loader_batch]] = loss1[n1 + loader_batch]
314 for n2 in range(input2.size(0)):
315 losses_I[index_I[n2]] = loss2[n2]
316
317 losses_V_aug1_slt = (losses_V_aug1 - losses_V_aug1.min()) / (losses_V_aug1.max() - losses_V_aug1.min())
318 losses_V_aug2_slt = (losses_V_aug2 - losses_V_aug2.min()) / (losses_V_aug2.max() - losses_V_aug2.min())
319 losses_I_slt = (losses_I - losses_I.min()) / (losses_I.max() - losses_I.min())
320 losses_V_slt = torch.cat((losses_V_aug1_slt, losses_V_aug2_slt), 0)
321
322 input_loss_V = losses_V_slt.reshape(-1, 1)
323 input_loss_I = losses_I_slt.reshape(-1, 1)
324
325 # fit a two-component GMM to the loss
326 gmm_V = GaussianMixture(n_components=2, max_iter=100, tol=1e-2, reg_covar=5e-4)
327 gmm_V.fit(input_loss_V)
328 prob_V = gmm_V.predict_proba(input_loss_V)
329 prob_V = prob_V[:, gmm_V.means_.argmin()]
330
331 gmm_I = GaussianMixture(n_components=2, max_iter=100, tol=1e-2, reg_covar=5e-4)
332 gmm_I.fit(input_loss_I)
333 prob_I = gmm_I.predict_proba(input_loss_I)
334 prob_I = prob_I[:, gmm_I.means_.argmin()]
335
336 return prob_V, prob_I
337
338
339def train(epoch, net, optimizer, trainloader):

Callers 1

run.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected