use pyo3::prelude::*;
use crate::python_interface::modules::PySequential;
use crate::python_interface::layers::PyEmbeddingBag;
use crate::nn::policy::Policy;
#[pyclass(name="Policy")]
pub struct PyPolicy {
pub policy: Box<Policy>,
}
#[pymethods]
impl PyPolicy {
#[new]
pub fn new(embeddings: PyEmbeddingBag, common: PySequential, action_net: PySequential, value_net: PySequential, obs_perms: Vec<Vec<usize>>, act_perms: Vec<Vec<usize>>) -> Self {
let policy = Box::new(Policy::new(embeddings.embedding, common.seq, action_net.seq, value_net.seq, obs_perms, act_perms));
PyPolicy { policy }
}
pub fn predict(&self, obs: Vec<usize>, masks: Vec<bool>) -> (Vec<f32>, f32) {
self.policy.predict(obs, masks)
}
pub fn forward(&self, obs: Vec<usize>, masks: Vec<bool>) -> (Vec<f32>, f32) {
self.policy.forward(obs, masks)
}
pub fn full_predict(&self, obs: Vec<usize>, masks: Vec<bool>) -> (Vec<f32>, f32) {
self.policy.full_predict(obs, masks)
}
}