| 58 | } |
| 59 | |
| 60 | void Classifier::Train(int num_epochs, int batch_size, float learning_rate, std::string train_val_dir, std::string image_type, std::string save_path){ |
| 61 | std::string path_train = train_val_dir+ "\\train"; |
| 62 | std::string path_val = train_val_dir + "\\val"; |
| 63 | |
| 64 | auto custom_dataset_train = dataSetClc(path_train, image_type).map(torch::data::transforms::Stack<>()); |
| 65 | auto custom_dataset_val = dataSetClc(path_val, image_type).map(torch::data::transforms::Stack<>()); |
| 66 | |
| 67 | auto data_loader_train = torch::data::make_data_loader<torch::data::samplers::RandomSampler>(std::move(custom_dataset_train), batch_size); |
| 68 | auto data_loader_val = torch::data::make_data_loader<torch::data::samplers::RandomSampler>(std::move(custom_dataset_val), batch_size); |
| 69 | |
| 70 | float loss_train = 0; float loss_val = 0; |
| 71 | float acc_train = 0.0; float acc_val = 0.0; float best_acc = 0.0; |
| 72 | for (size_t epoch = 1; epoch <= num_epochs; ++epoch) { |
| 73 | size_t batch_index_train = 0; |
| 74 | size_t batch_index_val = 0; |
| 75 | if (epoch == int(num_epochs / 2)) { learning_rate /= 10; } |
| 76 | torch::optim::Adam optimizer(vgg->parameters(), learning_rate); // Learning Rate |
| 77 | if (epoch < int(num_epochs / 8)) |
| 78 | { |
| 79 | for (auto mm : vgg->named_parameters()) |
| 80 | { |
| 81 | if (strstr(mm.key().data(), "classifier")) |
| 82 | { |
| 83 | mm.value().set_requires_grad(true); |
| 84 | } |
| 85 | else |
| 86 | { |
| 87 | mm.value().set_requires_grad(false); |
| 88 | } |
| 89 | } |
| 90 | } |
| 91 | else { |
| 92 | for (auto mm : vgg->named_parameters()) |
| 93 | { |
| 94 | mm.value().set_requires_grad(true); |
| 95 | } |
| 96 | } |
| 97 | // Iterate data loader to yield batches from the dataset |
| 98 | for (auto& batch : *data_loader_train) { |
| 99 | auto data = batch.data; |
| 100 | auto target = batch.target.squeeze(); |
| 101 | data = data.to(torch::kF32).to(device).div(255.0); |
| 102 | target = target.to(torch::kInt64).to(device); |
| 103 | optimizer.zero_grad(); |
| 104 | // Execute the model |
| 105 | torch::Tensor prediction = vgg->forward(data); |
| 106 | //cout << prediction << endl; |
| 107 | auto acc = prediction.argmax(1).eq(target).sum(); |
| 108 | acc_train += acc.template item<float>() / batch_size; |
| 109 | // Compute loss value |
| 110 | torch::Tensor loss = torch::nll_loss(prediction, target); |
| 111 | // Compute gradients |
| 112 | loss.backward(); |
| 113 | // Update the parameters |
| 114 | optimizer.step(); |
| 115 | loss_train += loss.item<float>(); |
| 116 | batch_index_train++; |
| 117 | std::cout << "Epoch: " << epoch << " |Train Loss: " << loss_train / batch_index_train << " |Train Acc:" << acc_train / batch_index_train << "\r"; |