MCPcopy Create free account
hub / github.com/AllentDan/LibtorchTutorials / Train

Method Train

lesson5-TrainingVGG/Classification.cpp:60–146  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

58}
59
60void 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";

Callers 1

mainFunction · 0.45

Calls 5

dataSetClcClass · 0.85
dataMethod · 0.80
keyMethod · 0.45
valueMethod · 0.45
forwardMethod · 0.45

Tested by

no test coverage detected