From c6569733a1d77e466b34a3aaa10cbe047dfa8eb0 Mon Sep 17 00:00:00 2001 From: Christian van der Loo Date: Sun, 14 May 2023 16:15:36 -0400 Subject: [PATCH] implemented basic RL --- Cargo.lock | 167 +++++++++++++++++++++++++++++++++++++--- Cargo.toml | 3 + src/agent/expectimax.rs | 6 +- src/agent/mod.rs | 5 +- src/agent/random.rs | 24 +++--- src/agent/rl.rs | 91 ++++++++++++++++++++++ src/agent/user.rs | 6 +- src/game.rs | 4 +- src/main.rs | 86 +++++++++++++++++++-- 9 files changed, 354 insertions(+), 38 deletions(-) create mode 100644 src/agent/rl.rs diff --git a/Cargo.lock b/Cargo.lock index 00d2ec8..9fe2dfd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -8,9 +8,12 @@ version = "0.1.0" dependencies = [ "crossterm 0.26.1", "enum-map", + "etcetera", "fastrand", "rayon", "rurel", + "serde", + "serde_json", "strum", "strum_macros", "transpose", @@ -148,7 +151,18 @@ checksum = "2a4da76b3b6116d758c7ba93f7ec6a35d2e2cf24feda76c6e38a375f4d5c59f2" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 1.0.109", +] + +[[package]] +name = "etcetera" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "136d1b5283a1ab77bd9257427ffd09d8667ced0570b6f938942bc7568ed5b943" +dependencies = [ + "cfg-if", + "home", + "windows-sys 0.48.0", ] [[package]] @@ -186,6 +200,15 @@ dependencies = [ "libc", ] +[[package]] +name = "home" +version = "0.5.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5444c27eef6923071f7ebcc33e3444508466a76f7a2b93da00ed6e19f30c1ddb" +dependencies = [ + "windows-sys 0.48.0", +] + [[package]] name = "instant" version = "0.1.12" @@ -195,6 +218,12 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "itoa" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "453ad9f582a441959e5f0d088b02ce04cfe8d51a8eaf077f12ac6d3e94164ca6" + [[package]] name = "libc" version = "0.2.142" @@ -238,7 +267,7 @@ dependencies = [ "libc", "log", "wasi", - "windows-sys", + "windows-sys 0.45.0", ] [[package]] @@ -290,7 +319,7 @@ dependencies = [ "libc", "redox_syscall", "smallvec", - "windows-sys", + "windows-sys 0.45.0", ] [[package]] @@ -393,12 +422,49 @@ version = "1.0.12" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4f3208ce4d8448b3f3e7d168a73f5e0c43a61e32930de3bceeccedb388b6bf06" +[[package]] +name = "ryu" +version = "1.0.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f91339c0467de62360649f8d3e185ca8de4224ff281f66000de5eb2a77a79041" + [[package]] name = "scopeguard" version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d29ab0c6d3fc0ee92fe66e2d99f700eab17a8d57d1c1d3b748380fb20baa78cd" +[[package]] +name = "serde" +version = "1.0.163" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2113ab51b87a539ae008b5c6c02dc020ffa39afd2d83cffcb3f4eb2722cebec2" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.163" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8c805777e3930c8883389c602315a24224bcc738b63905ef87cd1420353ea93e" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.16", +] + +[[package]] +name = "serde_json" +version = "1.0.96" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "057d394a50403bcac12672b2b18fb387ab6d289d957dab67dd201875391e52f1" +dependencies = [ + "itoa", + "ryu", + "serde", +] + [[package]] name = "signal-hook" version = "0.3.15" @@ -457,7 +523,7 @@ dependencies = [ "proc-macro2", "quote", "rustversion", - "syn", + "syn 1.0.109", ] [[package]] @@ -471,6 +537,17 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "syn" +version = "2.0.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6f671d4b5ffdb8eadec19c0ae67fe2639df8684bd7bc4b83d986b8db549cf01" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + [[package]] name = "transpose" version = "0.2.2" @@ -546,7 +623,16 @@ version = "0.45.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "75283be5efb2831d37ea142365f009c02ec203cd29a3ebecbc093d52315b66d0" dependencies = [ - "windows-targets", + "windows-targets 0.42.2", +] + +[[package]] +name = "windows-sys" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "677d2418bec65e3338edb076e806bc1ec15693c5d0104683f2efe857f61056a9" +dependencies = [ + "windows-targets 0.48.0", ] [[package]] @@ -555,13 +641,28 @@ version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8e5180c00cd44c9b1c88adb3693291f1cd93605ded80c250a75d472756b4d071" dependencies = [ - "windows_aarch64_gnullvm", - "windows_aarch64_msvc", - "windows_i686_gnu", - "windows_i686_msvc", - "windows_x86_64_gnu", - "windows_x86_64_gnullvm", - "windows_x86_64_msvc", + "windows_aarch64_gnullvm 0.42.2", + "windows_aarch64_msvc 0.42.2", + "windows_i686_gnu 0.42.2", + "windows_i686_msvc 0.42.2", + "windows_x86_64_gnu 0.42.2", + "windows_x86_64_gnullvm 0.42.2", + "windows_x86_64_msvc 0.42.2", +] + +[[package]] +name = "windows-targets" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b1eb6f0cd7c80c79759c929114ef071b87354ce476d9d94271031c0497adfd5" +dependencies = [ + "windows_aarch64_gnullvm 0.48.0", + "windows_aarch64_msvc 0.48.0", + "windows_i686_gnu 0.48.0", + "windows_i686_msvc 0.48.0", + "windows_x86_64_gnu 0.48.0", + "windows_x86_64_gnullvm 0.48.0", + "windows_x86_64_msvc 0.48.0", ] [[package]] @@ -570,38 +671,80 @@ version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "597a5118570b68bc08d8d59125332c54f1ba9d9adeedeef5b99b02ba2b0698f8" +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91ae572e1b79dba883e0d315474df7305d12f569b400fcf90581b06062f7e1bc" + [[package]] name = "windows_aarch64_msvc" version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e08e8864a60f06ef0d0ff4ba04124db8b0fb3be5776a5cd47641e942e58c4d43" +[[package]] +name = "windows_aarch64_msvc" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b2ef27e0d7bdfcfc7b868b317c1d32c641a6fe4629c171b8928c7b08d98d7cf3" + [[package]] name = "windows_i686_gnu" version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c61d927d8da41da96a81f029489353e68739737d3beca43145c8afec9a31a84f" +[[package]] +name = "windows_i686_gnu" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "622a1962a7db830d6fd0a69683c80a18fda201879f0f447f065a3b7467daa241" + [[package]] name = "windows_i686_msvc" version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "44d840b6ec649f480a41c8d80f9c65108b92d89345dd94027bfe06ac444d1060" +[[package]] +name = "windows_i686_msvc" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4542c6e364ce21bf45d69fdd2a8e455fa38d316158cfd43b3ac1c5b1b19f8e00" + [[package]] name = "windows_x86_64_gnu" version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8de912b8b8feb55c064867cf047dda097f92d51efad5b491dfb98f6bbb70cb36" +[[package]] +name = "windows_x86_64_gnu" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ca2b8a661f7628cbd23440e50b05d705db3686f894fc9580820623656af974b1" + [[package]] name = "windows_x86_64_gnullvm" version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "26d41b46a36d453748aedef1486d5c7a85db22e56aff34643984ea85514e94a3" +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7896dbc1f41e08872e9d5e8f8baa8fdd2677f29468c4e156210174edc7f7b953" + [[package]] name = "windows_x86_64_msvc" version = "0.42.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9aec5da331524158c6d1a4ac0ab1541149c0b9505fde06423b02f5ef0106b9f0" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1a515f5799fe4961cb532f983ce2b23082366b898e52ffbce459c86f67c8378a" diff --git a/Cargo.toml b/Cargo.toml index fc06673..2471849 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -8,9 +8,12 @@ edition = "2021" [dependencies] crossterm = "0.26.1" enum-map = "2.5.0" +etcetera = "0.8.0" fastrand = "1.9.0" rayon = "1.7.0" rurel = "0.4.0" +serde = { version = "1.0.163", features = ["derive"] } +serde_json = "1.0.96" strum = "0.24.1" strum_macros = "0.24.3" transpose = "0.2.2" diff --git a/src/agent/expectimax.rs b/src/agent/expectimax.rs index ce7a7f6..697f7a4 100644 --- a/src/agent/expectimax.rs +++ b/src/agent/expectimax.rs @@ -11,8 +11,8 @@ pub struct Expectimax { scores: MoveScores, } -impl Agent for Expectimax { - fn new(game: Game) -> Self +impl Expectimax { + pub fn new(game: Game) -> Self where Self: Sized, { @@ -22,7 +22,9 @@ impl Agent for Expectimax { scores: MoveScores::default(), } } +} +impl Agent for Expectimax { fn next_move(&mut self) { todo!() } diff --git a/src/agent/mod.rs b/src/agent/mod.rs index 97212b8..f04c578 100644 --- a/src/agent/mod.rs +++ b/src/agent/mod.rs @@ -6,12 +6,9 @@ use crate::{game::Game, IntAction}; pub mod expectimax; pub mod random; pub mod user; +pub mod rl; pub trait Agent { - fn new(game: Game) -> Self - where - Self: Sized; - fn get_game(&self) -> &Game; fn get_input(&mut self, _: &Event) -> IntAction { IntAction::Continue diff --git a/src/agent/random.rs b/src/agent/random.rs index d3068d0..5c78a40 100644 --- a/src/agent/random.rs +++ b/src/agent/random.rs @@ -12,10 +12,13 @@ pub struct RandomAgent { game: Game, } -impl Agent for RandomAgent { - fn new(game: Game) -> Self { +impl RandomAgent { + pub fn new(game: Game) -> Self { RandomAgent { game } } +} + +impl Agent for RandomAgent { fn next_move(&mut self) { self.game.update(match fastrand::usize(0..4) { @@ -61,15 +64,7 @@ pub enum RandomTreeMetric { } impl RandomTree { - pub fn new_with(game: Game, metric: RandomTreeMetric) -> Self { - let mut ag = RandomTree::new(game); - ag.metric = metric; - ag - } -} - -impl Agent for RandomTree { - fn new(game: Game) -> Self { + pub fn new(game: Game) -> Self { RandomTree { game, sim_count: 1000, @@ -77,7 +72,14 @@ impl Agent for RandomTree { last_scores: MoveScores::default(), } } + pub fn new_with(game: Game, metric: RandomTreeMetric) -> Self { + let mut ag = RandomTree::new(game); + ag.metric = metric; + ag + } +} +impl Agent for RandomTree { fn next_move(&mut self) { let mut scores = MoveScores::default(); for game_move in Move::iter() { diff --git a/src/agent/rl.rs b/src/agent/rl.rs new file mode 100644 index 0000000..2830ff3 --- /dev/null +++ b/src/agent/rl.rs @@ -0,0 +1,91 @@ +use etcetera::{choose_base_strategy, BaseStrategy}; +use rurel::{ + mdp::{Agent, State}, + AgentTrainer, +}; +use serde::{Deserialize, Serialize}; +use std::{collections::HashMap, fs}; +use strum::IntoEnumIterator; + +use crate::game::{Game, Move}; + +use super::Agent as GameAgent; + +pub const STRATEGY: BaseStrategy = choose_base_strategy().unwrap(); +pub const STORE_PATH: &str = STRATEGY.data_dir(); + +#[derive(PartialEq, Eq, Hash, Clone, Debug)] +pub struct RLAgent { + game: Game, +} + +impl RLAgent { + pub fn new(game: Game) -> Self + where + Self: Sized, + { + RLAgent { game } + } +} + +impl State for Game { + type A = Move; + + fn actions(&self) -> Vec { + Move::iter().collect() + } + + fn reward(&self) -> f64 { + if self.game_over() { + 0.0 + } else { + *self.get_score() as f64 + } + } +} + +impl Agent for RLAgent { + fn current_state(&self) -> &Game { + &self.game + } + + fn take_action(&mut self, action: &Move) { + self.game.update(*action); + } +} + +pub struct RLAgentTrained { + game: Game, + trainer: AgentTrainer, +} + +type TrainerExport = HashMap>; + +impl RLAgentTrained { + pub fn new(game: Game) -> Self + where + Self: Sized, + { + // read in trainer + let mut trainer = AgentTrainer::new(); + let data = fs::read_to_string(STORE_PATH).expect("Unable to read file"); + let imported_state: TrainerExport = serde_json::from_str(&data).unwrap(); + trainer.import_state(imported_state); + RLAgentTrained { game, trainer } + } +} + +impl GameAgent for RLAgentTrained { + fn next_move(&mut self) { + let Some(action) = self.trainer.best_action(&self.game) else { return; }; + self.game.update(action); + } + + fn get_game(&self) -> &Game { + &self.game + } + + fn messages(&self) -> Vec { + vec![tui::text::Spans::from("Performing RL actions.")] + } +} diff --git a/src/agent/user.rs b/src/agent/user.rs index b661927..d797fb7 100644 --- a/src/agent/user.rs +++ b/src/agent/user.rs @@ -13,14 +13,16 @@ pub struct UserAgent { game: Game, } -impl Agent for UserAgent { - fn new(game: Game) -> Self +impl UserAgent { + pub fn new(game: Game) -> Self where Self: Sized, { UserAgent { game } } +} +impl Agent for UserAgent { fn next_move(&mut self) { thread::sleep(Duration::from_millis(10)); } diff --git a/src/game.rs b/src/game.rs index 88be8eb..0485d7d 100644 --- a/src/game.rs +++ b/src/game.rs @@ -1,8 +1,10 @@ use core::fmt; use enum_map::Enum; +use serde::{Deserialize, Serialize}; use strum_macros::EnumIter; +#[derive(Debug, PartialEq, Eq, Hash, Serialize, Deserialize)] pub struct Game { state: [u16; 16], score: usize, @@ -32,7 +34,7 @@ impl Clone for Game { } } -#[derive(Enum, EnumIter, Debug, PartialEq, Clone, Copy)] +#[derive(Enum, EnumIter, Debug, PartialEq, Eq, Hash, Clone, Copy, Serialize, Deserialize)] pub enum Move { Up, Down, diff --git a/src/main.rs b/src/main.rs index 6e6c2aa..d732769 100644 --- a/src/main.rs +++ b/src/main.rs @@ -2,12 +2,19 @@ use crate::agent::{random::RandomAgent, Agent}; use crate::game::*; use agent::random::{RandomTree, RandomTreeMetric}; +use agent::rl::{RLAgent, RLAgentTrained, STORE_PATH}; use agent::user::UserAgent; use crossterm::{ event::{self, DisableMouseCapture, EnableMouseCapture, Event, KeyCode}, execute, terminal::{disable_raw_mode, enable_raw_mode, EnterAlternateScreen, LeaveAlternateScreen}, }; +use rurel::strategy::explore::RandomExploration; +use rurel::strategy::learn::QLearning; +use rurel::strategy::terminate::FixedIterations; +use rurel::AgentTrainer; +use std::fs::{File}; +use std::io::Write; use std::sync::RwLock; use std::thread::JoinHandle; use std::{error::Error, io, sync::Arc, thread, time::Duration}; @@ -29,6 +36,8 @@ static MENU_ITEMS: &[&str] = &[ "Solve (Random)", "Solve (Tree Search, Max Score)", "Solve (Tree Search, Max Moves)", + "Train (RL)", + "Solve (RL)", "Solve (Expectimax)", ]; @@ -37,6 +46,7 @@ pub enum Screen { state: ListState, menu: List<'static>, }, + Train(JoinHandle<()>), // join handle for multithreading if needed Game(JoinHandle<()>, Arc>>), } @@ -155,6 +165,15 @@ fn ui(f: &mut Frame, app: &mut App) { let paragraph = Paragraph::new(text).block(block).wrap(Wrap { trim: true }); f.render_widget(paragraph, chunks[1]); } + Screen::Train(_) => { + let block = Block::default().title("Training").borders(Borders::ALL); + let text = vec![ + Spans::from("Training in progress..."), + Spans::from("Press q to exit"), + ]; + let paragraph = Paragraph::new(text).block(block).wrap(Wrap { trim: true }); + f.render_widget(paragraph, chunks[0]); + } Screen::Game(_, game_sim) => { let agent = game_sim.read().unwrap(); let game = agent.get_game(); @@ -192,6 +211,12 @@ pub enum IntAction { Exit, } +pub enum MenuItem { + Play(Box), + Train, + Exit, +} + fn get_interaction(app: &mut App, timeout: Duration) -> Result { // each tick, lets see what screen we're at for interaction match &mut app.screen { @@ -226,14 +251,43 @@ fn get_interaction(app: &mut App, timeout: Duration) -> Result { let game = Game::new(); - let agent: Box = match state.selected() { - Some(0) => Box::new(UserAgent::new(game)), - Some(1) => Box::new(RandomAgent::new(game)), - Some(2) => Box::new(RandomTree::new(game)), - Some(3) => Box::new(RandomTree::new_with(game, RandomTreeMetric::AvgMoves)), + let item: MenuItem = match state.selected() { + Some(0) => MenuItem::Play(Box::new(UserAgent::new(game))), + Some(1) => MenuItem::Play(Box::new(RandomAgent::new(game))), + Some(2) => MenuItem::Play(Box::new(RandomTree::new(game))), + Some(3) => MenuItem::Play(Box::new(RandomTree::new_with( + game, + RandomTreeMetric::AvgMoves, + ))), + Some(4) => MenuItem::Train, + Some(5) => MenuItem::Play(Box::new(RLAgentTrained::new(game))), _ => panic!(), }; + let MenuItem::Play(agent) = item else { + match item { + MenuItem::Train => { + let t = thread::spawn(move || { + let mut trainer = AgentTrainer::new(); + let mut agent = RLAgent::new(Game::new()); + trainer.train(&mut agent, &QLearning::new(0.2, 0.01, 2.), &mut FixedIterations::new(1000000), &RandomExploration::new()); + // write out to file + let mut file = File::create(STORE_PATH).unwrap(); + let Ok(res) = serde_json::to_string(&trainer.export_learned_values()) else { + return; + }; + file.write_all(res.as_bytes()).unwrap(); + }); + app.screen = Screen::Train(t); + } + MenuItem::Exit => { + return Ok(IntAction::Exit); + } + _ => {} + }; + return Ok(IntAction::Continue); + }; + let agent = Arc::new(RwLock::new(agent)); let local_agent = agent.clone(); let t = thread::spawn(move || { @@ -241,11 +295,31 @@ fn get_interaction(app: &mut App, timeout: Duration) -> Result {} }; } + Screen::Train(t) => { + if t.is_finished() { + return Ok(IntAction::Exit); + } + + if !event::poll(timeout)? { + return Ok(IntAction::Continue); + } + let event = event::read()?; + let Event::Key(key_event) = event else { + return Ok(IntAction::Continue); + }; + match key_event.code { + KeyCode::Char('q') => { + return Ok(IntAction::Exit); + } + _ => {} + }; + } Screen::Game(_, agent) => { if !event::poll(timeout)? { return Ok(IntAction::Continue); @@ -283,7 +357,7 @@ fn run_tui(terminal: &mut Terminal) -> io::Result<()> { IntAction::Continue => {} IntAction::Exit => match app.screen { Screen::Menu { state: _, menu: _ } => break, - Screen::Game(_, _) => { + Screen::Game(_, _) | Screen::Train(_) => { app.screen = Screen::default(); continue; } -- 2.51.2