https://mooseframework.inl.gov
GaussianProcessTrainer.h
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 #ifdef MOOSE_LIBTORCH_ENABLED
10 
11 #pragma once
12 
13 #include "SurrogateTrainer.h"
14 #include "Standardizer.h"
15 
16 #include "Distribution.h"
17 
18 #include "CovarianceFunctionBase.h"
19 #include "CovarianceInterface.h"
20 
21 #include "GaussianProcess.h"
22 
23 #include "LibtorchUtils.h"
24 
26 {
27 public:
30  virtual void preTrain() override;
31  virtual void train() override;
32  virtual void postTrain() override;
33 
35  const StochasticTools::GaussianProcess & gp() const { return _gp; }
36 
37 private:
39  const std::vector<Real> & _predictor_row;
40 
43 
45  std::vector<std::vector<Real>> _params_buffer;
46 
48  std::vector<std::vector<Real>> _data_buffer;
49 
51  torch::Tensor & _training_params;
52 
54  torch::Tensor _training_data;
55 
58 
61 
63  bool _do_tuning;
64 
67 
69  const std::vector<Real> & _sampler_row;
70 };
71 
72 #endif
const StochasticTools::GaussianProcess::GPOptimizerOptions _optimization_opts
Struct holding parameters necessary for parameter tuning.
const std::vector< Real > & _sampler_row
Data from the current sampler row.
virtual void train() override
const StochasticTools::GaussianProcess & gp() const
const InputParameters & parameters() const
torch::Tensor _training_data
Data (y) used for training.
bool _do_tuning
Flag to toggle hyperparameter tuning/optimization.
Structure containing the optimization options for hyperparameter-tuning.
virtual void postTrain() override
StochasticTools::GaussianProcess & gp()
GaussianProcessTrainer(const InputParameters &parameters)
torch::Tensor & _training_params
Paramaters (x) used for training, along with statistics.
static InputParameters validParams()
const std::vector< Real > & _predictor_row
Data from the current predictor row.
virtual void preTrain() override
This is the main trainer base class.
bool _standardize_data
Switch for training data(y) standardization.
std::vector< std::vector< Real > > _data_buffer
Data (y) used for training.
StochasticTools::GaussianProcess & _gp
Gaussian process handler responsible for managing training related tasks.
Utility class dedicated to hold structures and functions commont to Gaussian Processes.
std::vector< std::vector< Real > > _params_buffer
Parameters (x) used for training – we&#39;ll allgather these in postTrain().
bool _standardize_params
Switch for training param (x) standardization.