add dino jumping environment, deterministic/reproducable behaviour and save-and-load feature

- add feature to save and load learning progress (Q-Table) and current episode count
- episode end is now purely decided by environment instead of monte carlo algo capping it on 10 actions
- using linkedHashMap on all locations to ensure deterministic behaviour
- fixed major RNG issue to reproduce algorithmic behaviour
- clearing rewardHistory, to only save the last 10k rewards
- added google dino jump environment
This commit is contained in:
2019-12-22 23:33:56 +01:00
parent b1246f62cc
commit 5a4e380faf
24 changed files with 415 additions and 56 deletions
@@ -0,0 +1,13 @@
package evironment.jumpingDino;
public class Config {
public static final int FRAME_WIDTH = 1280;
public static final int FRAME_HEIGHT = 720;
public static final int GROUND_Y = 50;
public static final int DINO_STARTING_X = 50;
public static final int DINO_SIZE = 50;
public static final int OBSTACLE_SIZE = 60;
public static final int OBSTACLE_SPEED = 30;
public static final int DINO_JUMP_SPEED = 20;
public static final int MAX_JUMP_HEIGHT = 200;
}
@@ -0,0 +1,40 @@
package evironment.jumpingDino;
import lombok.Getter;
import java.awt.*;
public class Dino extends RenderObject {
@Getter
private boolean inJump;
public Dino(int size, int x, int y, int dx, int dy, Color color) {
super(size, x, y, dx, dy, color);
}
public void jump(){
if(!inJump){
dy = -Config.DINO_JUMP_SPEED;
inJump = true;
}
}
private void fall(){
if(inJump){
dy = Config.DINO_JUMP_SPEED;
}
}
@Override
public void tick(){
// reached max jump height
if(y + dy < Config.FRAME_HEIGHT - Config.GROUND_Y -Config.OBSTACLE_SIZE - Config.MAX_JUMP_HEIGHT){
fall();
}else if(y + dy >= Config.FRAME_HEIGHT - Config.GROUND_Y - Config.DINO_SIZE){
inJump = false;
dy = 0;
y = Config.FRAME_HEIGHT - Config.GROUND_Y - Config.DINO_SIZE;
}
super.tick();
}
}
@@ -0,0 +1,6 @@
package evironment.jumpingDino;
public enum DinoAction {
JUMP,
NOTHING,
}
@@ -0,0 +1,32 @@
package evironment.jumpingDino;
import core.State;
import lombok.AllArgsConstructor;
import lombok.Getter;
import java.io.Serializable;
@AllArgsConstructor
@Getter
public class DinoState implements State, Serializable {
private int xDistanceToObstacle;
@Override
public String toString() {
return Integer.toString(xDistanceToObstacle);
}
@Override
public int hashCode() {
return this.xDistanceToObstacle;
}
@Override
public boolean equals(Object obj) {
if(obj instanceof DinoState){
DinoState toCompare = (DinoState) obj;
return toCompare.getXDistanceToObstacle() == this.xDistanceToObstacle;
}
return super.equals(obj);
}
}
@@ -0,0 +1,78 @@
package evironment.jumpingDino;
import core.Environment;
import core.State;
import core.StepResultEnvironment;
import core.gui.Visualizable;
import evironment.jumpingDino.gui.DinoWorldComponent;
import lombok.Getter;
import javax.swing.*;
import java.awt.*;
@Getter
public class DinoWorld implements Environment<DinoAction>, Visualizable {
private Dino dino;
private Obstacle currentObstacle;
public DinoWorld(){
dino = new Dino(Config.DINO_SIZE, Config.DINO_STARTING_X, Config.FRAME_HEIGHT - Config.GROUND_Y - Config.DINO_SIZE, 0, 0, Color.GREEN);
spawnNewObstacle();
}
private boolean ranIntoObstacle(){
Obstacle o = currentObstacle;
Dino p = dino;
boolean xAxis = (o.getX() <= p.getX() && p.getX() < o.getX() + Config.OBSTACLE_SIZE)
|| (o.getX() <= p.getX() + Config.DINO_SIZE && p.getX() + Config.DINO_SIZE < o.getX() + Config.OBSTACLE_SIZE);
boolean yAxis = (o.getY() <= p.getY() && p.getY() < o.getY() + Config.OBSTACLE_SIZE)
|| (o.getY() <= p.getY() + Config.DINO_SIZE && p.getY() + Config.DINO_SIZE < o.getY() + Config.OBSTACLE_SIZE);
return xAxis && yAxis;
}
private int getDistanceToObstacle(){
return currentObstacle.getX() - dino.getX() + Config.DINO_SIZE;
}
@Override
public StepResultEnvironment step(DinoAction action) {
boolean done = false;
int reward = 1;
if(action == DinoAction.JUMP){
dino.jump();
}
dino.tick();
currentObstacle.tick();
if(currentObstacle.getX() < -Config.OBSTACLE_SIZE){
spawnNewObstacle();
}
if(ranIntoObstacle()){
done = true;
}
return new StepResultEnvironment(new DinoState(getDistanceToObstacle()), reward, done, "");
}
private void spawnNewObstacle(){
currentObstacle = new Obstacle(Config.OBSTACLE_SIZE, Config.FRAME_WIDTH + Config.OBSTACLE_SIZE, Config.FRAME_HEIGHT - Config.GROUND_Y - Config.OBSTACLE_SIZE, -Config.OBSTACLE_SPEED, 0, Color.BLACK);
}
private void spawnDino(){
dino = new Dino(Config.DINO_SIZE, Config.DINO_STARTING_X, Config.FRAME_HEIGHT - Config.GROUND_Y - Config.DINO_SIZE, 0, 0, Color.GREEN);
}
@Override
public State reset() {
spawnDino();
spawnNewObstacle();
return new DinoState(getDistanceToObstacle());
}
@Override
public JComponent visualize() {
return new DinoWorldComponent(this);
}
}
@@ -0,0 +1,10 @@
package evironment.jumpingDino;
import java.awt.*;
public class Obstacle extends RenderObject {
public Obstacle(int size, int x, int y, int dx, int dy, Color color) {
super(size, x, y, dx, dy, color);
}
}
@@ -0,0 +1,28 @@
package evironment.jumpingDino;
import lombok.AllArgsConstructor;
import lombok.Getter;
import java.awt.*;
@AllArgsConstructor
@Getter
public abstract class RenderObject {
protected int size;
protected int x;
protected int y;
protected int dx;
protected int dy;
protected Color color;
public void render(Graphics g){
g.setColor(color);
g.fillRect(x, y, size, size);
}
public void tick(){
y += dy;
x += dx;
}
}
@@ -0,0 +1,27 @@
package evironment.jumpingDino.gui;
import evironment.jumpingDino.Config;
import evironment.jumpingDino.DinoWorld;
import javax.swing.*;
import java.awt.*;
public class DinoWorldComponent extends JComponent {
private DinoWorld dinoWorld;
public DinoWorldComponent(DinoWorld dinoWorld){
this.dinoWorld = dinoWorld;
setPreferredSize(new Dimension(Config.FRAME_WIDTH, Config.FRAME_HEIGHT));
setVisible(true);
}
@Override
protected void paintComponent(Graphics g) {
super.paintComponent(g);
g.setColor(Color.BLACK);
g.fillRect(0, Config.FRAME_HEIGHT - Config.GROUND_Y, Config.FRAME_WIDTH, 2);
dinoWorld.getDino().render(g);
dinoWorld.getCurrentObstacle().render(g);
}
}