12 #ifdef MOOSE_LIBTORCH_ENABLED 24 exportHyperParameter(
const torch::Tensor & tensor)
27 if (cpu_tensor.scalar_type() != at::kDouble)
28 cpu_tensor = cpu_tensor.to(at::kDouble).contiguous();
29 const auto flattened = cpu_tensor.reshape({-1});
30 return {flattened.data_ptr<
Real>(), flattened.data_ptr<Real>() + flattened.numel()};
34 exportScalarHyperParameter(
const torch::Tensor & tensor)
37 if (cpu_tensor.scalar_type() != at::kDouble)
38 cpu_tensor = cpu_tensor.to(at::kDouble);
39 return cpu_tensor.item<
Real>();
50 "Tool for extracting hyperparameter data from gaussian process user object and " 51 "storing in VectorPostprocessor vectors.");
52 params.
addRequiredParam<UserObjectName>(
"gp_name",
"Name of GaussianProcess.");
68 for (
const auto & iter : hyperparam_map)
73 _hp_vector.back()->push_back(exportScalarHyperParameter(iter.second));
78 mooseError(
"Unsupported hyperparameter rank ", iter.second.dim(),
" for ", iter.first,
".");
80 const auto vec = exportHyperParameter(iter.second);
81 for (
unsigned int ii = 0; ii < vec.size(); ++ii)
static bool isVectorHyperParameter(const torch::Tensor &tensor)
Return true if a hyperparameter tensor stores a vector of values.
const StochasticTools::GaussianProcess & getGP() const
torch::Tensor toCPUContiguous(const torch::Tensor &tensor)
static bool isScalarHyperParameter(const torch::Tensor &tensor)
Return true if a hyperparameter tensor stores one scalar value.
virtual void initialize() override
static InputParameters validParams()
static InputParameters validParams()
const GaussianProcessSurrogate & _gp_surrogate
Reference to GaussianProcess.
VectorPostprocessorValue & declareVector(const std::string &vector_name)
std::vector< VectorPostprocessorValue * > _hp_vector
Vector of hyperparamater values.
GaussianProcessData(const InputParameters ¶meters)
DIE A HORRIBLE DEATH HERE typedef LIBMESH_DEFAULT_SCALAR_TYPE Real
Interface for objects that need to use samplers.
void mooseError(Args &&... args) const
registerMooseObject("StochasticToolsApp", GaussianProcessData)
static InputParameters validParams()