Skip to main content

ValueNetwork

Trait ValueNetwork 

Source
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§

Source

fn evaluate_value(&self, observations: &Array2<T>) -> Result<Array1<T>>

Evaluate value function for given observations

Source

fn update_parameters( &mut self, deltas: &HashMap<String, Array1<T>>, ) -> Result<()>

Add a parameter delta to the value-function parameters.

Source

fn get_parameters(&self) -> HashMap<String, Array1<T>>

Get current value function parameters

Provided Methods§

Source

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".

Implementors§

Source§

impl<T: Float + Debug + Send + Sync + 'static> ValueNetwork<T> for LinearQFunction<T>

Source§

impl<T: Float + Debug + Send + Sync + 'static> ValueNetwork<T> for LinearValueFunction<T>