()
| 11 | cuda_exist = code_flag == 0 |
| 12 | |
| 13 | def 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()) |
no test coverage detected