https://mooseframework.inl.gov
Loading...
Searching...
No Matches
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
19#include "CovarianceInterface.h"
20
21#include "GaussianProcess.h"
22
23#include "LibtorchUtils.h"
24
26{
27public:
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
37private:
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
64
67
69 const std::vector<Real> & _sampler_row;
70};
71
72#endif
const StochasticTools::GaussianProcess & gp() const
virtual void postTrain() override
StochasticTools::GaussianProcess & gp()
virtual void preTrain() override
StochasticTools::GaussianProcess & _gp
Gaussian process handler responsible for managing training related tasks.
std::vector< std::vector< Real > > _data_buffer
Data (y) used for training.
bool _do_tuning
Flag to toggle hyperparameter tuning/optimization.
torch::Tensor _training_data
Data (y) used for training.
bool _standardize_data
Switch for training data(y) standardization.
const std::vector< Real > & _predictor_row
Data from the current predictor row.
std::vector< std::vector< Real > > _params_buffer
Parameters (x) used for training – we'll allgather these in postTrain().
const StochasticTools::GaussianProcess::GPOptimizerOptions _optimization_opts
Struct holding parameters necessary for parameter tuning.
static InputParameters validParams()
const std::vector< Real > & _sampler_row
Data from the current sampler row.
torch::Tensor & _training_params
Paramaters (x) used for training, along with statistics.
virtual void train() override
bool _standardize_params
Switch for training param (x) standardization.
const InputParameters & parameters() const
Utility class dedicated to hold structures and functions commont to Gaussian Processes.
This is the main trainer base class.
Structure containing the optimization options for hyperparameter-tuning.