use anyhow::Result;
use serde::{Deserialize, Serialize};
use std::{
fs::File,
io::{BufReader, Write},
path::Path,
};
#[derive(Debug, Deserialize, Serialize, PartialEq, Clone)]
pub struct TrainerConfig {
pub(super) max_opts: usize,
pub(super) eval_episodes: usize,
pub(super) eval_threshold: Option<f32>,
pub(super) model_dir: Option<String>,
pub(super) opt_interval: usize,
pub(super) eval_interval: usize,
pub(super) record_interval: usize,
pub(super) save_interval: usize,
}
impl Default for TrainerConfig {
fn default() -> Self {
Self {
max_opts: 0,
eval_interval: 0,
eval_episodes: 0,
eval_threshold: None,
model_dir: None,
opt_interval: 1,
record_interval: usize::MAX,
save_interval: usize::MAX,
}
}
}
impl TrainerConfig {
pub fn max_opts(mut self, v: usize) -> Self {
self.max_opts = v;
self
}
pub fn eval_interval(mut self, v: usize) -> Self {
self.eval_interval = v;
self
}
pub fn eval_episodes(mut self, v: usize) -> Self {
self.eval_episodes = v;
self
}
pub fn eval_threshold(mut self, v: f32) -> Self {
self.eval_threshold = Some(v);
self
}
pub fn model_dir<T: Into<String>>(mut self, model_dir: T) -> Self {
self.model_dir = Some(model_dir.into());
self
}
pub fn opt_interval(mut self, opt_interval: usize) -> Self {
self.opt_interval = opt_interval;
self
}
pub fn record_interval(mut self, record_interval: usize) -> Self {
self.record_interval = record_interval;
self
}
pub fn save_interval(mut self, save_interval: usize) -> Self {
self.save_interval = save_interval;
self
}
pub fn load(path: impl AsRef<Path>) -> Result<Self> {
let file = File::open(path)?;
let rdr = BufReader::new(file);
let b = serde_yaml::from_reader(rdr)?;
Ok(b)
}
pub fn save(&self, path: impl AsRef<Path>) -> Result<()> {
let mut file = File::create(path)?;
file.write_all(serde_yaml::to_string(&self)?.as_bytes())?;
Ok(())
}
}