pub mod normalizer;
use std::{collections::BTreeMap, fmt::Debug, sync::Arc};
use serde::{Deserialize, Serialize};
use crate::error::{Error, InvalidParameterError};
use crate::tensor::R2lTensor;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub enum Space<T: R2lTensor> {
Discrete(usize),
Box {
min: Option<T>,
max: Option<T>,
shape: Vec<usize>,
},
MultiDiscrete {
nvec: T,
shape: Vec<usize>,
},
MultiBinary {
shape: Vec<usize>,
},
Tuple(Vec<Space<T>>),
Dict(BTreeMap<String, Space<T>>),
}
impl<T: R2lTensor> Space<T> {
pub fn convert<U: R2lTensor>(&self) -> Result<Space<U>, Error> {
Ok(match self {
Self::Discrete(size) => Space::Discrete(*size),
Self::Box { min, max, shape } => Space::Box {
min: min.as_ref().map(U::convert).transpose()?,
max: max.as_ref().map(U::convert).transpose()?,
shape: shape.clone(),
},
Self::MultiDiscrete { nvec, shape } => Space::MultiDiscrete {
nvec: U::convert(nvec)?,
shape: shape.clone(),
},
Self::MultiBinary { shape } => Space::MultiBinary {
shape: shape.clone(),
},
Self::Tuple(spaces) => {
Space::Tuple(spaces.iter().map(Self::convert).collect::<Result<_, _>>()?)
}
Self::Dict(spaces) => Space::Dict(
spaces
.iter()
.map(|(key, space)| Ok((key.clone(), space.convert()?)))
.collect::<Result<_, Error>>()?,
),
})
}
pub fn shape(&self) -> Option<&[usize]> {
match self {
Self::Discrete(_) => Some(&[]),
Self::Box { shape, .. }
| Self::MultiDiscrete { shape, .. }
| Self::MultiBinary { shape } => Some(shape),
Self::Tuple(_) | Self::Dict(_) => None,
}
}
pub fn size(&self) -> usize {
match &self {
Self::Discrete(size) => *size,
Self::Box { shape, .. }
| Self::MultiDiscrete { shape, .. }
| Self::MultiBinary { shape, .. } => shape.iter().product(),
Self::Tuple(spaces) => spaces.iter().map(Self::size).sum(),
Self::Dict(spaces) => spaces.values().map(Self::size).sum(),
}
}
#[must_use]
pub fn action_size(&self) -> usize {
match self {
Self::Discrete(_) => 1,
Self::Box { shape, .. }
| Self::MultiDiscrete { shape, .. }
| Self::MultiBinary { shape } => shape.iter().product(),
Self::Tuple(spaces) => spaces.iter().map(Self::action_size).sum(),
Self::Dict(spaces) => spaces.values().map(Self::action_size).sum(),
}
}
}
#[derive(Debug, Clone)]
pub struct EnvDescription<T: R2lTensor> {
pub observation_space: Space<T>,
pub action_space: Space<T>,
}
impl<T: R2lTensor> EnvDescription<T> {
pub fn new(observation_space: Space<T>, action_space: Space<T>) -> Self {
Self {
observation_space,
action_space,
}
}
pub fn action_size(&self) -> usize {
self.action_space.action_size()
}
pub fn observation_size(&self) -> usize {
self.observation_space.size()
}
}
pub struct Snapshot<T: R2lTensor> {
pub state: T,
pub reward: f32,
pub terminated: bool,
pub truncated: bool,
}
impl<T: R2lTensor> Snapshot<T> {
pub fn new(state: T, reward: f32, terminated: bool, truncated: bool) -> Self {
Self {
state,
reward,
terminated,
truncated,
}
}
pub fn done(&self) -> bool {
self.terminated || self.truncated
}
}
pub trait Env {
type Tensor: R2lTensor;
fn reset(&mut self, seed: u64) -> Result<Self::Tensor, Error>;
fn step(&mut self, action: Self::Tensor) -> Result<Snapshot<Self::Tensor>, Error>;
fn env_description(&self) -> EnvDescription<Self::Tensor>;
}
pub trait EnvBuilder: Sync + Send + 'static {
type Env: Env;
fn build_env(&self) -> Result<Self::Env, Error>;
fn env_description(&self) -> Result<EnvDescription<<Self::Env as Env>::Tensor>, Error> {
let env = self.build_env()?;
Ok(env.env_description())
}
}
impl<E: Env, F: Sync + Send + 'static> EnvBuilder for F
where
F: Fn() -> Result<E, Error>,
{
type Env = E;
fn build_env(&self) -> Result<E, Error> {
(self)()
}
}
pub struct EnvBuilderType<EB: EnvBuilder>(EnvBuilderKind<EB>);
enum EnvBuilderKind<EB: EnvBuilder> {
Homogeneous {
builder: Arc<EB>,
n_envs: usize,
},
Heterogeneous {
builders: Vec<Arc<EB>>,
},
}
impl<EB: EnvBuilder> Clone for EnvBuilderType<EB> {
fn clone(&self) -> Self {
Self(match &self.0 {
EnvBuilderKind::Homogeneous { builder, n_envs } => EnvBuilderKind::Homogeneous {
builder: builder.clone(),
n_envs: *n_envs,
},
EnvBuilderKind::Heterogeneous { builders } => EnvBuilderKind::Heterogeneous {
builders: builders.clone(),
},
})
}
}
impl<EB: EnvBuilder> EnvBuilderType<EB> {
fn from_kind(kind: EnvBuilderKind<EB>) -> Result<Self, Error> {
match &kind {
EnvBuilderKind::Homogeneous { n_envs: 0, .. } => {
return Err(Error::InvalidParameter(Box::new(
InvalidParameterError::InvalidValue {
name: "n_envs".into(),
expected: "a value greater than zero".into(),
value: "0".into(),
},
)));
}
EnvBuilderKind::Heterogeneous { builders } if builders.is_empty() => {
return Err(Error::InvalidParameter(Box::new(
InvalidParameterError::InvalidValue {
name: "builders".into(),
expected: "at least one environment builder".into(),
value: "empty".into(),
},
)));
}
_ => {}
}
Ok(Self(kind))
}
pub fn homogeneous(builder: EB, n_envs: usize) -> Result<Self, Error> {
Self::from_kind(EnvBuilderKind::Homogeneous {
builder: Arc::new(builder),
n_envs,
})
}
pub fn heterogeneous(builders: Vec<EB>) -> Result<Self, Error> {
Self::from_kind(EnvBuilderKind::Heterogeneous {
builders: builders.into_iter().map(Arc::new).collect(),
})
}
pub fn build_idx(&self, idx: usize) -> Result<EB::Env, Error> {
let n_envs = self.num_envs();
if idx >= self.num_envs() {
return Err(Error::InvalidParameter(Box::new(
InvalidParameterError::InvalidValue {
name: "environment index".into(),
expected: format!("an index below {n_envs}"),
value: idx.to_string(),
},
)));
}
match &self.0 {
EnvBuilderKind::Homogeneous { builder, .. } => builder.build_env(),
EnvBuilderKind::Heterogeneous { builders } => builders[idx].build_env(),
}
}
#[must_use]
pub fn num_envs(&self) -> usize {
match &self.0 {
EnvBuilderKind::Homogeneous { n_envs, .. } => *n_envs,
EnvBuilderKind::Heterogeneous { builders } => builders.len(),
}
}
pub fn env_description(&self) -> Result<EnvDescription<<EB::Env as Env>::Tensor>, Error> {
match &self.0 {
EnvBuilderKind::Homogeneous { builder, n_envs: _ } => builder.env_description(),
EnvBuilderKind::Heterogeneous { builders } => builders[0].env_description(),
}
}
}
pub fn action_ranges(nvec: &[usize]) -> impl Iterator<Item = (usize, usize)> + '_ {
nvec.iter().scan(0, |offset, choices| {
let start = *offset;
*offset += *choices;
Some((start, *choices))
})
}