add results for convergence for advanced dino jumping

This commit is contained in:
2020-03-05 13:17:54 +01:00
parent e67f40ad65
commit 4641f50b79
10 changed files with 58 additions and 8 deletions
@@ -107,7 +107,7 @@ public abstract class EpisodicLearning<A extends Enum> extends Learning<A> imple
if(timestampCurrentEpisode > 300000){
converged = true;
// t
File file = new File("convergence.txt");
File file = new File("convergenceAdv.txt");
try {
Files.writeString(Path.of(file.getPath()), currentEpisode/2 + ",", StandardOpenOption.APPEND);
} catch (IOException e) {
@@ -10,6 +10,9 @@ import lombok.Getter;
import javax.swing.*;
import java.awt.*;
/**
* 57 states
*/
@Getter
public class DinoWorld implements Environment<DinoAction>, Visualizable {
protected Dino dino;
@@ -1,9 +1,19 @@
package evironment.jumpingDino;
import core.RNG;
import core.State;
import java.awt.*;
/**
* 3580 states
* if:
* dx = -(int)((Math.random() + 0.5) * Config.OBSTACLE_SPEED);
* xSpawn = Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE;
*
* 350 states
* if 4 speed variants
*/
public class DinoWorldAdvanced extends DinoWorld{
public DinoWorldAdvanced(){
super();
@@ -18,9 +28,18 @@ public class DinoWorldAdvanced extends DinoWorld{
protected void spawnNewObstacle() {
int dx;
int xSpawn;
dx = -(int)((Math.random() + 0.5) * Config.OBSTACLE_SPEED);
double ran = RNG.getRandom().nextDouble();
if(ran < 0.25){
dx = -(int)(0.7 * Config.OBSTACLE_SPEED);
}else if(ran < 0.5){
dx = -(int)(1.3 * Config.OBSTACLE_SPEED);
}else if(ran < 0.75){
dx = -(int)(1.6 * Config.OBSTACLE_SPEED);
} else{
dx = -2 * Config.OBSTACLE_SPEED;
}
// randomly spawning more right outside of the screen
xSpawn = (int)(Math.random() + 0.5 * Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE);
xSpawn = Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE;
currentObstacle = new Obstacle(Config.OBSTACLE_SIZE, xSpawn, Config.FRAME_HEIGHT - Config.GROUND_Y - Config.OBSTACLE_SIZE, dx, 0, Color.BLACK);
}
}
+4 -3
View File
@@ -5,6 +5,7 @@ import core.algo.Method;
import core.controller.RLController;
import evironment.jumpingDino.DinoAction;
import evironment.jumpingDino.DinoWorld;
import evironment.jumpingDino.DinoWorldAdvanced;
import java.io.File;
import java.io.IOException;
@@ -14,7 +15,7 @@ import java.nio.file.StandardOpenOption;
public class DinoSampling {
public static void main(String[] args) {
File file = new File("convergence.txt");
File file = new File("convergenceAdv.txt");
for(float f = 0.05f; f <=1.003 ; f+=0.05f){
try {
Files.writeString(Path.of(file.getPath()), f + ",", StandardOpenOption.APPEND);
@@ -26,14 +27,14 @@ public class DinoSampling {
RNG.setSeed(i *13);
RLController<DinoAction> rl = new RLController<>(
new DinoWorld(),
new DinoWorldAdvanced(),
Method.MC_CONTROL_FIRST_VISIT,
DinoAction.values());
rl.setDelay(0);
rl.setDiscountFactor(1f);
rl.setEpsilon(f);
rl.setLearningRate(1f);
rl.setNrOfEpisodes(20000);
rl.setNrOfEpisodes(100000);
rl.start();
}
+3 -2
View File
@@ -6,13 +6,14 @@ import core.controller.RLController;
import core.controller.RLControllerGUI;
import evironment.jumpingDino.DinoAction;
import evironment.jumpingDino.DinoWorld;
import evironment.jumpingDino.DinoWorldAdvanced;
public class JumpingDino {
public static void main(String[] args) {
RNG.setSeed(55);
RLController<DinoAction> rl = new RLControllerGUI<>(
new DinoWorld(),
new DinoWorldAdvanced(),
Method.MC_CONTROL_FIRST_VISIT,
DinoAction.values());
@@ -20,7 +21,7 @@ public class JumpingDino {
rl.setDiscountFactor(1f);
rl.setEpsilon(0.15f);
rl.setLearningRate(1f);
rl.setNrOfEpisodes(10000);
rl.setNrOfEpisodes(1000000);
rl.start();
}
}