use super::types::EpisodeInfo;
use async_trait::async_trait;
use rlmesh_adapters::v1::{CustomTransform, EncodingTransform, ResolvedAdapter, Value};
use crate::spaces::{EnvContract, SpaceSpec, SpaceValue};
use crate::{Result, model::types::ModelObservation};
pub trait PredictFn: Send + Sync {
fn predict(&self, model_input: Value, episode: Option<&EpisodeInfo>) -> Result<Value>;
fn predict_chunk(
&self,
_model_input: Value,
_execution_horizon: u32,
_episode: Option<&EpisodeInfo>,
) -> Result<Option<Value>> {
Ok(None)
}
fn has_chunk(&self) -> bool {
false
}
fn predict_batch(&self, _inputs: Vec<Value>, _episodes: &[EpisodeInfo]) -> Result<Vec<Value>> {
Err(crate::Error::model("predict_batch is not implemented"))
}
fn has_batch(&self) -> bool {
false
}
fn predict_chunk_batch(
&self,
_inputs: Vec<Value>,
_execution_horizon: u32,
_episodes: &[EpisodeInfo],
) -> Result<Vec<Value>> {
Err(crate::Error::model(
"predict_chunk_batch is not implemented",
))
}
fn has_chunk_batch(&self) -> bool {
false
}
fn predict_spec_less(&self, observation: ModelObservation) -> Result<Vec<SpaceValue>>;
fn predict_spec_less_chunked(
&self,
observation: ModelObservation,
_execution_horizon: u32,
) -> Result<super::handler::PredictFrames> {
Ok(super::handler::PredictFrames {
actions: self.predict_spec_less(observation)?,
replay: Vec::new(),
})
}
fn allow_fusion(&self) -> bool {
false
}
fn native_chunk(&self) -> Option<u32> {
None
}
fn on_episode_end(&self, _episode_id: &str) -> Result<()> {
Ok(())
}
fn on_close(&self) -> Result<()> {
Ok(())
}
}
pub struct RouteConfig {
pub(crate) adapter: ResolvedAdapter,
pub(crate) observation_space: SpaceSpec,
pub(crate) action_space: SpaceSpec,
pub(crate) customs: Box<dyn CustomTransform + Send + Sync>,
pub(crate) encodings: Box<dyn EncodingTransform + Send + Sync>,
pub(crate) execution_horizon: u32,
pub(crate) delivers_history: bool,
}
impl RouteConfig {
pub fn new(
adapter: ResolvedAdapter,
observation_space: SpaceSpec,
action_space: SpaceSpec,
customs: Box<dyn CustomTransform + Send + Sync>,
encodings: Box<dyn EncodingTransform + Send + Sync>,
) -> Self {
Self {
adapter,
observation_space,
action_space,
customs,
encodings,
execution_horizon: 1,
delivers_history: false,
}
}
}
#[async_trait]
pub trait RouteResolver: Send + Sync {
async fn resolve(
&self,
route_key: &str,
env_contract: &EnvContract,
) -> Result<Option<RouteConfig>>;
}