diff --git a/src/game.rs b/src/game.rs index 0485d7d..2c8e1c9 100644 --- a/src/game.rs +++ b/src/game.rs @@ -4,9 +4,9 @@ use enum_map::Enum; use serde::{Deserialize, Serialize}; use strum_macros::EnumIter; -#[derive(Debug, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[derive(Debug, PartialEq, Eq, Hash, Clone, Copy, Serialize, Deserialize)] pub struct Game { - state: [u16; 16], + state: [u8; 16], score: usize, num_moves: usize, } @@ -24,16 +24,6 @@ impl Default for Game { } } -impl Clone for Game { - fn clone(&self) -> Self { - Game { - state: self.state, - score: self.score, - num_moves: self.num_moves, - } - } -} - #[derive(Enum, EnumIter, Debug, PartialEq, Eq, Hash, Clone, Copy, Serialize, Deserialize)] pub enum Move { Up, @@ -53,7 +43,12 @@ impl fmt::Display for Move { } } -fn merge_duplicates(v: &Vec, mut f: F) -> Vec +/// This function generates a new tile in a random empty spot. The new tile +/// will be a 2 with 90% probability and a 4 with 10% probability. +/// +/// * `v`: the current state of the game +/// * `f`: a function to call when a new tile is generated +fn merge_duplicates(v: &Vec, mut f: F) -> Vec where F: FnMut(usize) -> (), { @@ -81,6 +76,12 @@ impl Game { Game::default() } + pub fn new_from(state: [u8; 16]) -> Self { + let mut game = Game::default(); + game.state = state; + game + } + pub fn update(&mut self, input: Move) -> bool { let state_before = self.state.clone(); self.shift(input); @@ -97,11 +98,17 @@ impl Game { (x + y * 4) as usize } - pub fn get_tile(&self, x: u8, y: u8) -> u16 { + pub fn is_available_move(&self, input: Move) -> bool { + let mut g = self.clone(); + g.shift(input); + g.state != self.state + } + + pub fn get_tile(&self, x: u8, y: u8) -> u8 { self.state[Game::xy_to_index(x, y)] } - pub fn set_tile(&mut self, x: u8, y: u8, value: u16) { + pub fn set_tile(&mut self, x: u8, y: u8, value: u8) { self.state[Game::xy_to_index(x, y)] = value; } @@ -146,15 +153,15 @@ impl Game { &self.score } - pub fn get_state(&self) -> &[u16; 16] { + pub fn get_state(&self) -> &[u8; 16] { &self.state } - pub fn set_state(&mut self, s: [u16; 16]) { + pub fn set_state(&mut self, s: [u8; 16]) { self.state = s; } - fn get_condensed_rows(&self) -> Vec> { + fn get_condensed_rows(&self) -> Vec> { // get vec of rows with empty removed self.state .chunks(4) @@ -162,7 +169,7 @@ impl Game { .collect() } - fn get_condensed_cols(&self) -> Vec> { + fn get_condensed_cols(&self) -> Vec> { // get vec of cols with empty removed self.state .iter() @@ -181,7 +188,7 @@ impl Game { Move::Left | Move::Right => self.get_condensed_rows(), }; - let new_state: [u16; 16] = condensed + let new_state: [u8; 16] = condensed .iter() .map(|v| merge_duplicates(v, |s| self.score += s)) .map(|mut v| { @@ -210,7 +217,7 @@ impl Game { } fn generate_tile(&mut self) { - let n = fastrand::u16(1..=2); + let n = fastrand::u8(1..=2); // get indexes of empty tiles let empty_indexes = self diff --git a/src/main.rs b/src/main.rs index d732769..a51f9a1 100644 --- a/src/main.rs +++ b/src/main.rs @@ -2,7 +2,7 @@ use crate::agent::{random::RandomAgent, Agent}; use crate::game::*; use agent::random::{RandomTree, RandomTreeMetric}; -use agent::rl::{RLAgent, RLAgentTrained, STORE_PATH}; +use agent::rl::{get_trainer, RLAgent, RLAgentTrained}; use agent::user::UserAgent; use crossterm::{ event::{self, DisableMouseCapture, EnableMouseCapture, Event, KeyCode}, @@ -11,9 +11,9 @@ use crossterm::{ }; use rurel::strategy::explore::RandomExploration; use rurel::strategy::learn::QLearning; -use rurel::strategy::terminate::FixedIterations; +use rurel::strategy::terminate::SinkStates; use rurel::AgentTrainer; -use std::fs::{File}; +use std::fs::File; use std::io::Write; use std::sync::RwLock; use std::thread::JoinHandle; @@ -160,6 +160,7 @@ fn ui(f: &mut Frame, app: &mut App) { let block = Block::default().title("Info").borders(Borders::ALL); let text = vec![ Spans::from("Use arrow keys to navigate"), + Spans::from(format!("Writing to {}", agent::rl::data_file_path())), Spans::from("Press q to exit"), ]; let paragraph = Paragraph::new(text).block(block).wrap(Wrap { trim: true }); @@ -268,14 +269,14 @@ fn get_interaction(app: &mut App, timeout: Duration) -> Result { 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()); + let mut trainer = get_trainer(); + for _ in 0..10000 { + let mut agent = RLAgent::new(Game::new()); + trainer.train(&mut agent, &QLearning::new(0.2, 0.01, 2.), &mut SinkStates {}, &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; - }; + let mut file = File::create(agent::rl::data_file_path()).unwrap(); + let res = ron::to_string(&trainer.export_learned_values()).unwrap(); file.write_all(res.as_bytes()).unwrap(); }); app.screen = Screen::Train(t);