Following the api of other layers

This commit is contained in:
himanshupathak21061998
2020-06-12 03:28:36 +05:30
parent c6499e0b12
commit 968d38d332
2 changed files with 25 additions and 1 deletions
@@ -79,6 +79,21 @@ class RBF
OutputDataType const& OutputParameter() const { return outputParameter; }
//! Modify the output parameter.
OutputDataType& OutputParameter() { return outputParameter; }
//! Get the parameters.
OutputDataType const& Parameters() const { return weights; }
//! Modify the parameters.
OutputDataType& Parameters() { return weights; }
//! Get the input parameter.
InputDataType const& InputParameter() const { return inputParameter; }
//! Modify the input parameter.
InputDataType& InputParameter() { return inputParameter; }
//! Get the input size.
size_t InputSize() const { return inSize; }
//! Get the output size.
size_t OutputSize() const { return outSize; }
//! Get the detla.
OutputDataType const& Delta() const { return delta; }
@@ -104,12 +119,16 @@ class RBF
//! Locally-stored the learnable scaling factor of the shape.
InputDataType sigmas;
//! Locally-stored the outeput distances of the shape.
//! Locally-stored input parameter object.
InputDataType inputParameter;
//! Locally-stored the output distances of the shape.
OutputDataType distances;
//! Locally-stored weight object.
OutputDataType weights;
//! Locally-stored reset parameter used to initialize the layer once.
bool reset;
//! Locally-stored number of input units.
@@ -89,6 +89,11 @@ void RBF<InputDataType, OutputDataType>::serialize(
ar & BOOST_SERIALIZATION_NVP(distances);
ar & BOOST_SERIALIZATION_NVP(sigmas);
ar & BOOST_SERIALIZATION_NVP(centres);
// This is inefficient, but we have to allocate this memory so that
// WeightSetVisitor gets the right size.
if (Archive::is_loading::value)
weights.set_size(outSize * inSize + outSize, 1);
}
} // namespace ann