27 const unsigned int num_batches)
31 if (num_samples < num_batches)
34 else if (num_samples % num_batches == 0)
35 return num_samples / num_batches;
41 const unsigned int sample_per_batch_1 = num_samples / num_batches;
42 const unsigned int remainder_1 = num_samples % num_batches;
43 const unsigned int sample_per_batch_2 = sample_per_batch_1 - 1;
44 const unsigned int remainder_2 =
45 num_samples - (num_samples / sample_per_batch_2) * sample_per_batch_2;
47 const Real rel_remainder1 = Real(remainder_1) / Real(sample_per_batch_1);
48 const Real rel_remainder2 = Real(remainder_2) / Real(sample_per_batch_2);
50 return rel_remainder2 > rel_remainder1 ? sample_per_batch_2 : sample_per_batch_1;
74 std::unique_ptr<torch::optim::Optimizer> optimizer;
78 optimizer = std::make_unique<torch::optim::Adam>(
79 nn.parameters(), torch::optim::AdamOptions(options.
learning_rate));
82 optimizer = std::make_unique<torch::optim::Adagrad>(nn.parameters(), options.
learning_rate);
85 optimizer = std::make_unique<torch::optim::RMSprop>(nn.parameters(), options.
learning_rate);
88 optimizer = std::make_unique<torch::optim::SGD>(nn.parameters(), options.
learning_rate);
101 const auto t_begin = MPI_Wtime();
111 int real_rank = processor_id();
113 int used_rank = real_rank < num_ranks ? real_rank : 0;
115 const auto num_samples = dataset.
size().value();
118 mooseError(
"The number of used processors* number of requestedf batches " +
120 " is greater than the number of samples used for the training!");
123 const unsigned int sample_per_batch = computeBatchSize(num_samples, options.
num_batches);
126 const unsigned int sample_per_proc = computeLocalBatchSize(sample_per_batch, num_ranks);
129 auto transformed_data_set = dataset.map(torch::data::transforms::Stack<>());
132 SamplerType sampler(num_samples, num_ranks, used_rank, options.
allow_duplicates);
136 torch::data::make_data_loader(std::move(transformed_data_set), sampler, sample_per_proc);
139 std::unique_ptr<torch::optim::Optimizer> optimizer = createOptimizer(_nn, options);
142 Real initial_loss = 1.0;
143 Real epoch_loss = 0.0;
146 unsigned int epoch = 1;
147 while (epoch <= options.num_epochs && rel_loss > options.
rel_loss_tol)
151 for (
auto & batch : *data_loader)
154 optimizer->zero_grad();
157 torch::Tensor prediction = _nn.forward(batch.data);
160 torch::Tensor loss = torch::mse_loss(prediction, batch.target);
166 if (real_rank == used_rank)
167 epoch_loss += loss.item<
double>();
174 for (
auto & param : _nn.named_parameters())
176 if (real_rank != used_rank)
177 param.value().grad().data() = param.value().grad().data() * 0.0;
179 MPI_Allreduce(MPI_IN_PLACE,
180 param.value().grad().data_ptr(),
181 param.value().grad().numel(),
184 _communicator.get());
186 param.value().grad().data() = param.value().grad().data() / num_ranks;
195 _communicator.sum(epoch_loss);
197 epoch_loss = epoch_loss / options.
num_batches / num_ranks;
200 initial_loss = epoch_loss;
202 rel_loss = epoch_loss / initial_loss;
207 Moose::out <<
"Epoch: " << epoch <<
" | Loss: " << COLOR_GREEN << epoch_loss
208 << COLOR_DEFAULT <<
" | Rel. loss: " << COLOR_GREEN << rel_loss << COLOR_DEFAULT
215 auto t_end = MPI_Wtime();
218 Moose::out <<
"Neural net training time: " << COLOR_GREEN << (t_end - t_begin) << COLOR_DEFAULT
219 <<
" s" << std::endl;