MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / train

Function train

CV/PaddleFSL/train.py:13–82  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

11cuda_exist = code_flag == 0
12
13def train():
14 args = parse_args()
15 print(args)
16
17 if args.use_gpu and cuda_exist:
18 place = fluid.CUDAPlace(0)
19 print('GPU is used...')
20 else:
21 place = fluid.CPUPlace()
22 print('CPU is used...')
23
24
25 with fluid.dygraph.guard(place):
26 print('start training ... ')
27
28 # prepare method
29 model = prepare_model(args)
30 model.train()
31
32 # prepare optimizer
33 opt = prepare_optimizer(args, model)
34
35 # prepare dataloader
36 train_data_batches, val_data_batches = prepare_dataloader(args)
37
38 save_name = args.method+'_'+args.backbone+'_'+str(args.k_shot)+'shot_'+str(args.n_way)+'way'
39 best_val_acc = 0
40 with LogWriter(logdir=args.log_dir+'logs/'+args.dataset+'/', filename_suffix='_'+save_name) as writer:
41 for epoch in range(args.epochs):
42 train_loss, train_acc = [], []
43 for batch_id, batch in enumerate(train_data_batches):
44 samples, label = batch
45 samples = fluid.dygraph.to_variable(samples)
46 labels = fluid.dygraph.to_variable(label)
47 loss, acc = model.loss(samples, labels)
48 avg_loss = fluid.layers.mean(loss)
49 train_loss.append(avg_loss.numpy())
50 train_acc.append(acc.numpy())
51
52 if batch_id % 100 == 0:
53 print("epoch: {}, batch_id: {}, loss is: {}, acc is: {}".format(epoch, batch_id, avg_loss.numpy(), acc.numpy()))
54 avg_loss.backward()
55 opt.minimize(avg_loss)
56 model.clear_gradients()
57
58 writer.add_scalar(tag="train_loss", step=epoch, value=np.mean(train_loss))
59 writer.add_scalar(tag="train_acc", step=epoch, value=np.mean(train_acc))
60
61 model.eval()
62 accuracies = []
63 losses = []
64 for batch_id, batch in enumerate(val_data_batches):
65 samples, label = batch
66 samples = fluid.dygraph.to_variable(samples)
67 labels = fluid.dygraph.to_variable(label)
68 loss, acc = model.loss(samples, labels)
69 avg_loss = fluid.layers.mean(loss)
70 accuracies.append(acc.numpy())

Callers 1

train.pyFile · 0.70

Calls 9

parse_argsFunction · 0.90
prepare_modelFunction · 0.90
prepare_optimizerFunction · 0.90
prepare_dataloaderFunction · 0.90
trainMethod · 0.45
lossMethod · 0.45
appendMethod · 0.45
evalMethod · 0.45
state_dictMethod · 0.45

Tested by

no test coverage detected