https://mooseframework.inl.gov
Loading...
Searching...
No Matches
LibtorchDRLControl.C
Go to the documentation of this file.
1//* This file is part of the MOOSE framework
2//* https://mooseframework.inl.gov
3//*
4//* All rights reserved, see COPYRIGHT for full restrictions
5//* https://github.com/idaholab/moose/blob/master/COPYRIGHT
6//*
7//* Licensed under LGPL 2.1, please see LICENSE for details
8//* https://www.gnu.org/licenses/lgpl-2.1.html
9
10#ifdef MOOSE_LIBTORCH_ENABLED
11
12#include "LibtorchDRLControl.h"
13#include "TorchScriptModule.h"
15#include "Transient.h"
16
18
21{
24 "Sets the value of multiple 'Real' input parameters and postprocessors based on a Deep "
25 "Reinforcement Learning (DRL) neural network trained using a PPO algorithm.");
26 params.addRequiredParam<std::vector<Real>>(
27 "action_standard_deviations", "Standard deviation value used while sampling the actions.");
28 params.addParam<unsigned int>("seed", "Seed for the random number generator.");
29
30 return params;
31}
32
34 : LibtorchNeuralNetControl(parameters),
35 _current_control_signal_log_probabilities(std::vector<Real>(_control_names.size(), 0.0)),
36 _action_std(getParam<std::vector<Real>>("action_standard_deviations"))
37{
38 if (_control_names.size() != _action_std.size())
39 paramError("action_standard_deviations",
40 "Number of action_standard_deviations does not match the number of controlled "
41 "parameters.");
42
43 // Fixing the RNG seed to make sure every experiment is the same.
44 if (isParamValid("seed"))
45 torch::manual_seed(getParam<unsigned int>("seed"));
46
47 // We convert and store the user-supplied standard deviations into a tensor which can be easily
48 // used by routines in libtorch
49 _std = torch::eye(_control_names.size());
50 for (unsigned int i = 0; i < _control_names.size(); ++i)
51 _std[i][i] = _action_std[i];
52}
53
54void
56{
57 if (_nn)
58 {
59 unsigned int n_controls = _control_names.size();
60 unsigned int num_old_timesteps = _input_timesteps - 1;
61
62 // Fill a vector with the current values of the responses
64
65 // If this is the first time this control is called and we need to use older values, fill up the
66 // needed old values using the initial values
67 if (_old_responses.empty())
68 _old_responses.assign(num_old_timesteps, _current_response);
69
70 // Organize the old an current solution into a tensor so we can evaluate the neural net
71 torch::Tensor input_tensor = prepareInputTensor();
72
73 // Evaluate the neural network to get the expected control value
74 torch::Tensor output_tensor = _nn->forward(input_tensor);
75
76 // Sample control value (action) from Gaussian distribution
77 torch::Tensor action = at::normal(output_tensor, _std);
78
79 // Compute log probability
80 torch::Tensor log_probability = computeLogProbability(action, output_tensor);
81
82 // Convert data
83 _current_control_signals = {action.data_ptr<Real>(), action.data_ptr<Real>() + action.size(1)};
84
85 _current_control_signal_log_probabilities = {log_probability.data_ptr<Real>(),
86 log_probability.data_ptr<Real>() +
87 log_probability.size(1)};
88
89 for (unsigned int control_i = 0; control_i < n_controls; ++control_i)
90 {
91 // We scale the controllable value for physically meaningful control action
92 setControllableValueByName<Real>(_control_names[control_i],
93 _current_control_signals[control_i] *
94 _action_scaling_factors[control_i]);
95 }
96
97 // We add the curent solution to the old solutions and move everything in there one step
98 // backward
99 std::rotate(_old_responses.rbegin(), _old_responses.rbegin() + 1, _old_responses.rend());
101 }
102}
103
104torch::Tensor
105LibtorchDRLControl::computeLogProbability(const torch::Tensor & action,
106 const torch::Tensor & output_tensor)
107{
108 // Logarithmic probability of taken action, given the current distribution.
109 torch::Tensor var = torch::matmul(_std, _std);
110
111 return -((action - output_tensor) * (action - output_tensor)) / (2.0 * var) - torch::log(_std) -
112 std::log(std::sqrt(2.0 * M_PI));
113}
114
115Real
116LibtorchDRLControl::getSignalLogProbability(const unsigned int signal_index) const
117{
118 mooseAssert(signal_index < _control_names.size(),
119 "The index of the requested control signal is not in the [0," +
120 std::to_string(_control_names.size()) + ") range!");
122}
123
124#endif
registerMooseObject("StochasticToolsApp", LibtorchDRLControl)
void addRequiredParam(const std::string &name, const std::string &doc_string)
void addParam(const std::string &name, const std::initializer_list< typename T::value_type > &value, const std::string &doc_string)
void addClassDescription(const std::string &doc_string)
A time-dependent, neural-network-based controller which is associated with a Proximal Policy Optimiza...
LibtorchDRLControl(const InputParameters &parameters)
Construct using input parameters.
const std::vector< Real > _action_std
Standard deviation for the actions, supplied by the user.
Real getSignalLogProbability(const unsigned int signal_index) const
Get the logarithmic probability of (signal_index)-th signal of the control neural net.
virtual void execute() override
We compute the actions in this function together with the corresponding logarithmic probabilities.
static InputParameters validParams()
torch::Tensor computeLogProbability(const torch::Tensor &action, const torch::Tensor &output_tensor)
Function which computes the logarithmic probability of given actions.
std::vector< Real > _current_control_signal_log_probabilities
The log probability of control signals from the last evaluation of the controller.
torch::Tensor _std
Standard deviations converted to a 2D diagonal tensor that can be used by Libtorch routines.
const std::vector< Real > _action_scaling_factors
const unsigned int _input_timesteps
std::vector< Real > _current_response
std::vector< std::vector< Real > > & _old_responses
static InputParameters validParams()
std::vector< Real > _current_control_signals
const std::vector< std::string > & _control_names
torch::Tensor prepareInputTensor()
std::shared_ptr< Moose::LibtorchNeuralNetBase > _nn
void paramError(const std::string &param, Args... args) const
bool isParamValid(const std::string &name) const