pub trait ValueNetwork<T: Float + Debug + Send + Sync + 'static> {
// Required methods
fn evaluate_value(&self, observations: &Array2<T>) -> Result<Array1<T>>;
fn update_parameters(
&mut self,
deltas: &HashMap<String, Array1<T>>,
) -> Result<()>;
fn get_parameters(&self) -> HashMap<String, Array1<T>>;
// Provided method
fn value_gradient(
&self,
observations: &Array2<T>,
residuals: &Array1<T>,
) -> Result<HashMap<String, Array1<T>>> { ... }
}Expand description
Value network interface for RL optimizers.
ValueNetwork::update_parameters follows the same additive delta
contract as PolicyNetwork::update_parameters.
Required Methods§
Sourcefn evaluate_value(&self, observations: &Array2<T>) -> Result<Array1<T>>
fn evaluate_value(&self, observations: &Array2<T>) -> Result<Array1<T>>
Evaluate value function for given observations
Sourcefn update_parameters(
&mut self,
deltas: &HashMap<String, Array1<T>>,
) -> Result<()>
fn update_parameters( &mut self, deltas: &HashMap<String, Array1<T>>, ) -> Result<()>
Add a parameter delta to the value-function parameters.
Sourcefn get_parameters(&self) -> HashMap<String, Array1<T>>
fn get_parameters(&self) -> HashMap<String, Array1<T>>
Get current value function parameters
Provided Methods§
Sourcefn value_gradient(
&self,
observations: &Array2<T>,
residuals: &Array1<T>,
) -> Result<HashMap<String, Array1<T>>>
fn value_gradient( &self, observations: &Array2<T>, residuals: &Array1<T>, ) -> Result<HashMap<String, Array1<T>>>
Gradient of a residual-weighted sum of value predictions:
∂/∂θ Σᵢ rᵢ · V(sᵢ).
Callers pass rᵢ = ∂L/∂V(sᵢ); for the mean-squared value loss
L = (1/N) Σ (V(sᵢ) − yᵢ)² that is rᵢ = 2(V(sᵢ) − yᵢ)/N.
Dyn Compatibility§
This trait is dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".