Compare commits
15
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ff6807dabd | ||
|
|
ad07c1da8f | ||
|
|
5b82e7965d | ||
|
|
4402d70467 | ||
|
|
b9be640284 | ||
|
|
a08b8160a3 | ||
|
|
595451e88b | ||
|
|
a40e279f48 | ||
|
|
9a3452ff9c | ||
|
|
740289ee2b | ||
|
|
e7404a8d24 | ||
|
|
0fde1bd962 | ||
|
|
f4b50627d1 | ||
|
|
78955a9521 | ||
|
|
328fc85214 |
@@ -17,5 +17,5 @@
|
|||||||
0.8,19870,3288,13724,4492,8159,5058,16764,5648,9462,19071,3914,1242,8262,26004,4036,9421,4914,2535,5362,7298,9587,37133,1837,35325,15272,14922,14138,7115,17236,5123,12157,37380,6086,37390,1672,15573,14241,2049,2602,6802,22362,7936,7544,5330,13155,16016,4544,1489,3780,6326,7794,31553,2808,1493,7788,12646,30464,22312,1681,12084,4163,2197,7950,22478,5106,26771,4382,10615,2586,12214,4799,6297,7589,4585,30365,32302,15734,5480,8626,7387,11932,4245,21532,1710,12737,7132,4740,14578,10680,8266,17300,4213,3264,35920,38026,10272,3984,2279,9739,33900
|
0.8,19870,3288,13724,4492,8159,5058,16764,5648,9462,19071,3914,1242,8262,26004,4036,9421,4914,2535,5362,7298,9587,37133,1837,35325,15272,14922,14138,7115,17236,5123,12157,37380,6086,37390,1672,15573,14241,2049,2602,6802,22362,7936,7544,5330,13155,16016,4544,1489,3780,6326,7794,31553,2808,1493,7788,12646,30464,22312,1681,12084,4163,2197,7950,22478,5106,26771,4382,10615,2586,12214,4799,6297,7589,4585,30365,32302,15734,5480,8626,7387,11932,4245,21532,1710,12737,7132,4740,14578,10680,8266,17300,4213,3264,35920,38026,10272,3984,2279,9739,33900
|
||||||
0.85,5493,10568,19366,5705,15430,8183,5721,13314,36667,33059,3753,40243,23888,25085,21843,6856,2803,9434,4794,29944,10730,39271,4484,23990,6350,16180,8099,4298,11220,4624,5946,24895,8464,4416,6619,2800,4081,12459,1981,12488,6380,9597,10328,1901,24563,13059,3639,12988,2604,4440,22666,1775,4078,5175,1144,3759,11119,1856,34970,10831,2229,5333,17121,9698,14919,2353,3963,8189,36145,13920,5301,16516,2446,46848,3985,,20640,151501,17556,1882,44216,39795,1638,57957,62050,3130,3693,5563,9780,3327,22969,39357,13749,37555,60070,9249,35426,4405,8340,18973
|
0.85,5493,10568,19366,5705,15430,8183,5721,13314,36667,33059,3753,40243,23888,25085,21843,6856,2803,9434,4794,29944,10730,39271,4484,23990,6350,16180,8099,4298,11220,4624,5946,24895,8464,4416,6619,2800,4081,12459,1981,12488,6380,9597,10328,1901,24563,13059,3639,12988,2604,4440,22666,1775,4078,5175,1144,3759,11119,1856,34970,10831,2229,5333,17121,9698,14919,2353,3963,8189,36145,13920,5301,16516,2446,46848,3985,,20640,151501,17556,1882,44216,39795,1638,57957,62050,3130,3693,5563,9780,3327,22969,39357,13749,37555,60070,9249,35426,4405,8340,18973
|
||||||
0.9,27355,24592,18962,2318,17604,35725,14327,38167,25602,50236,4999,9023,5562,7541,11799,25139,8724,12642,28509,57095,2147,5909,5414,12572,10018,68830,45393,18962,51656,25601,3444,45667,16813,57110,16492,3991,7315,17775,69277,34769,29824,11087,26371,3479,2540,9597,32593,13169,8588,2794,40136,56004,65307,24864,35523,19491,2673,5363,4799,5852,28566,42427,44011,40146,3757,1115,49574,5798,24249,2576,118943,6169,65584,7057,49505,116138,52083,1809,127776,3214,25689,103442,15260,62754,12390,3233,35309,68989,6615,30593,2503,29359,98237,11900,3240,64969,84134,25361,7384,13141
|
0.9,27355,24592,18962,2318,17604,35725,14327,38167,25602,50236,4999,9023,5562,7541,11799,25139,8724,12642,28509,57095,2147,5909,5414,12572,10018,68830,45393,18962,51656,25601,3444,45667,16813,57110,16492,3991,7315,17775,69277,34769,29824,11087,26371,3479,2540,9597,32593,13169,8588,2794,40136,56004,65307,24864,35523,19491,2673,5363,4799,5852,28566,42427,44011,40146,3757,1115,49574,5798,24249,2576,118943,6169,65584,7057,49505,116138,52083,1809,127776,3214,25689,103442,15260,62754,12390,3233,35309,68989,6615,30593,2503,29359,98237,11900,3240,64969,84134,25361,7384,13141
|
||||||
0.95,24269,14543,6828,3800,41079,47279,27177,17286,9802,7114,3756,85275,14507,34993,15139,15184,90742,27554,23713,6453,15157,7045,8048,47550,84540,93729,68601,6274,4713,30578,5024,94239,7315,8193,46871,96466,3695,70915,62947,32258,66228,2114,5084,12686,62905,19158,20940,36270,9037,34034,15016,15530,46276,11063,8586,15635,7196,70708,50836,22464,13463,86986,43541,2001,40565,28534,44700,5625,6552,16140,2450,8492,3304,22904,20951,100472,131147,131728,43674,514,79827,181148,31431,4761,1515,2075,138139,137795,71014170145,60000,42790,179835,18982,48085,28398,56788,126115,5442,118289,9386
|
0.95,24269,14543,6828,3800,41079,47279,27177,17286,9802,7114,3756,85275,14507,34993,15139,15184,90742,27554,23713,6453,15157,7045,8048,47550,84540,93729,68601,6274,4713,30578,5024,94239,7315,8193,46871,96466,3695,70915,62947,32258,66228,2114,5084,12686,62905,19158,20940,36270,9037,34034,15016,15530,46276,11063,8586,15635,7196,70708,50836,22464,13463,86986,43541,2001,40565,28534,44700,5625,6552,16140,2450,8492,3304,22904,20951,100472,131147,131728,43674,514,79827,181148,31431,4761,1515,2075,138139,137795,71014,170145,60000,42790,179835,18982,48085,28398,56788,126115,5442,118289
|
||||||
1.0,11364,6363,8012,109822,19730,8425,21388,7864,18427,34072,3126,52381,35105,86487,73913,88033,76264,105864,30103,9522,31049,3180,4838,4078,133687,39236,59239,22968,21540,98395,109063,4050,5612,4990,9933,83766,140114,116077,135653,130826,130070,92207,14994,87801,1577,70868,133816,79790,1587,23322,22071,13903,3584,9721,,38605,52375,67392,10075,97733,46173,29647,2558,28151,162569,4054,10537,30871,45538,97835,45132,35042,70203,3862,100614,84525,140691,81880,80914,35187,11596,51448,2945,56551,39236,84707,64324,100588,78645,12929,32701,63306,163991,2864,34802,72929,198161,71332,98627,137754
|
1.0,11364,6363,8012,109822,19730,8425,21388,7864,18427,34072,3126,52381,35105,86487,73913,88033,76264,105864,30103,9522,31049,3180,4838,4078,133687,39236,59239,22968,21540,98395,109063,4050,5612,4990,9933,83766,140114,116077,135653,130826,130070,92207,14994,87801,1577,70868,133816,79790,1587,23322,22071,13903,3584,9721,,38605,52375,67392,10075,97733,46173,29647,2558,28151,162569,4054,10537,30871,45538,97835,45132,35042,70203,3862,100614,84525,140691,81880,80914,35187,11596,51448,2945,56551,39236,84707,64324,100588,78645,12929,32701,63306,163991,2864,34802,72929,198161,71332,98627,137754
|
||||||
|
|||||||
@@ -17,24 +17,30 @@ public class RNG {
|
|||||||
private static Random rng;
|
private static Random rng;
|
||||||
private static Random rngEnv;
|
private static Random rngEnv;
|
||||||
private static int seed = 123;
|
private static int seed = 123;
|
||||||
|
private static int envSeed = 13;
|
||||||
static {
|
static {
|
||||||
rng = new Random();
|
rng = new Random();
|
||||||
|
rngEnv = new Random();
|
||||||
setSeed(seed, true);
|
setSeed(seed, true);
|
||||||
}
|
}
|
||||||
|
|
||||||
public static Random getRandom() {
|
public static Random getRandom() {
|
||||||
return rng;
|
return rng;
|
||||||
}
|
}
|
||||||
|
public static Random getRandomEnv() {
|
||||||
public static Random getEnvRandom() {
|
|
||||||
return rngEnv;
|
return rngEnv;
|
||||||
}
|
}
|
||||||
|
|
||||||
public static void setSeed(int seed, boolean setEnvSeed) {
|
public static void setSeed(int seed, boolean setEnvRandom) {
|
||||||
RNG.seed = seed;
|
RNG.seed = seed;
|
||||||
rng.setSeed(seed);
|
rng.setSeed(seed);
|
||||||
if(setEnvSeed) {
|
if(setEnvRandom) {
|
||||||
rngEnv.setSeed(seed);
|
rngEnv.setSeed(13);
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
public static void setSeed(int seed) {
|
||||||
|
setSeed(seed, true);
|
||||||
|
}
|
||||||
|
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -5,17 +5,12 @@ import core.Environment;
|
|||||||
import core.LearningConfig;
|
import core.LearningConfig;
|
||||||
import core.StepResult;
|
import core.StepResult;
|
||||||
import core.listener.LearningListener;
|
import core.listener.LearningListener;
|
||||||
import example.DinoSampling;
|
|
||||||
import lombok.Getter;
|
import lombok.Getter;
|
||||||
import lombok.Setter;
|
import lombok.Setter;
|
||||||
|
|
||||||
import java.io.File;
|
|
||||||
import java.io.IOException;
|
import java.io.IOException;
|
||||||
import java.io.ObjectInputStream;
|
import java.io.ObjectInputStream;
|
||||||
import java.io.ObjectOutputStream;
|
import java.io.ObjectOutputStream;
|
||||||
import java.nio.file.Files;
|
|
||||||
import java.nio.file.Path;
|
|
||||||
import java.nio.file.StandardOpenOption;
|
|
||||||
import java.util.ArrayList;
|
import java.util.ArrayList;
|
||||||
import java.util.List;
|
import java.util.List;
|
||||||
import java.util.concurrent.atomic.AtomicInteger;
|
import java.util.concurrent.atomic.AtomicInteger;
|
||||||
@@ -91,20 +86,6 @@ public abstract class EpisodicLearning<A extends Enum> extends Learning<A> imple
|
|||||||
super.dispatchStepEnd();
|
super.dispatchStepEnd();
|
||||||
timestamp++;
|
timestamp++;
|
||||||
timestampCurrentEpisode++;
|
timestampCurrentEpisode++;
|
||||||
// TODO: more sophisticated way to check convergence
|
|
||||||
if(timestampCurrentEpisode > 50000) {
|
|
||||||
converged = true;
|
|
||||||
// t
|
|
||||||
File file = new File(DinoSampling.FILE_NAME);
|
|
||||||
try {
|
|
||||||
Files.writeString(Path.of(file.getPath()), currentEpisode/2 + ",", StandardOpenOption.APPEND);
|
|
||||||
} catch (IOException e) {
|
|
||||||
e.printStackTrace();
|
|
||||||
}
|
|
||||||
System.out.println("converged after: " + currentEpisode/2 + " episode!");
|
|
||||||
episodesToLearn.set(0);
|
|
||||||
dispatchLearningEnd();
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
@Override
|
@Override
|
||||||
|
|||||||
@@ -16,8 +16,6 @@ import java.util.HashSet;
|
|||||||
import java.util.List;
|
import java.util.List;
|
||||||
import java.util.Set;
|
import java.util.Set;
|
||||||
import java.util.concurrent.CopyOnWriteArrayList;
|
import java.util.concurrent.CopyOnWriteArrayList;
|
||||||
import java.util.concurrent.ExecutorService;
|
|
||||||
import java.util.concurrent.Executors;
|
|
||||||
|
|
||||||
/**
|
/**
|
||||||
*
|
*
|
||||||
@@ -99,7 +97,7 @@ public abstract class Learning<A extends Enum>{
|
|||||||
|
|
||||||
public void save(ObjectOutputStream oos) throws IOException {
|
public void save(ObjectOutputStream oos) throws IOException {
|
||||||
oos.writeObject(rewardHistory);
|
oos.writeObject(rewardHistory);
|
||||||
oos.writeObject(stateActionTable);
|
// oos.writeObject(stateActionTable);
|
||||||
}
|
}
|
||||||
|
|
||||||
public void load(ObjectInputStream ois) throws IOException, ClassNotFoundException {
|
public void load(ObjectInputStream ois) throws IOException, ClassNotFoundException {
|
||||||
|
|||||||
@@ -3,8 +3,6 @@ package core.algo.mc;
|
|||||||
import core.*;
|
import core.*;
|
||||||
import core.algo.EpisodicLearning;
|
import core.algo.EpisodicLearning;
|
||||||
import core.policy.EpsilonGreedyPolicy;
|
import core.policy.EpsilonGreedyPolicy;
|
||||||
import core.policy.GreedyPolicy;
|
|
||||||
import core.policy.Policy;
|
|
||||||
import org.apache.commons.lang3.tuple.ImmutablePair;
|
import org.apache.commons.lang3.tuple.ImmutablePair;
|
||||||
import org.apache.commons.lang3.tuple.Pair;
|
import org.apache.commons.lang3.tuple.Pair;
|
||||||
|
|
||||||
@@ -14,7 +12,7 @@ import java.io.ObjectOutputStream;
|
|||||||
import java.util.*;
|
import java.util.*;
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Includes both variants of Monte-Carlo methods
|
* Includes both! variants of Monte-Carlo methods
|
||||||
* Default method is First-Visit.
|
* Default method is First-Visit.
|
||||||
* Change to Every-Visit by setting flag "useEveryVisit" in the constructor to true.
|
* Change to Every-Visit by setting flag "useEveryVisit" in the constructor to true.
|
||||||
* @param <A>
|
* @param <A>
|
||||||
@@ -25,17 +23,10 @@ public class MonteCarloControlEGreedy<A extends Enum> extends EpisodicLearning<A
|
|||||||
private Map<Pair<State, A>, Integer> returnCount;
|
private Map<Pair<State, A>, Integer> returnCount;
|
||||||
private boolean isEveryVisit;
|
private boolean isEveryVisit;
|
||||||
|
|
||||||
// t
|
|
||||||
private float epsilon;
|
|
||||||
// t
|
|
||||||
private Policy<A> greedyPolicy = new GreedyPolicy<>();
|
|
||||||
|
|
||||||
|
|
||||||
public MonteCarloControlEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, int delay, boolean useEveryVisit) {
|
public MonteCarloControlEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, int delay, boolean useEveryVisit) {
|
||||||
super(environment, actionSpace, discountFactor, delay);
|
super(environment, actionSpace, discountFactor, delay);
|
||||||
isEveryVisit = useEveryVisit;
|
isEveryVisit = useEveryVisit;
|
||||||
// t
|
|
||||||
this.epsilon = epsilon;
|
|
||||||
this.policy = new EpsilonGreedyPolicy<>(epsilon);
|
this.policy = new EpsilonGreedyPolicy<>(epsilon);
|
||||||
this.stateActionTable = new DeterministicStateActionTable<>(this.actionSpace);
|
this.stateActionTable = new DeterministicStateActionTable<>(this.actionSpace);
|
||||||
returnSum = new HashMap<>();
|
returnSum = new HashMap<>();
|
||||||
@@ -64,12 +55,7 @@ public class MonteCarloControlEGreedy<A extends Enum> extends EpisodicLearning<A
|
|||||||
|
|
||||||
while(envResult == null || !envResult.isDone()) {
|
while(envResult == null || !envResult.isDone()) {
|
||||||
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
||||||
A chosenAction;
|
A chosenAction = policy.chooseAction(actionValues);
|
||||||
if(currentEpisode % 2 == 1){
|
|
||||||
chosenAction = greedyPolicy.chooseAction(actionValues);
|
|
||||||
}else{
|
|
||||||
chosenAction = policy.chooseAction(actionValues);
|
|
||||||
}
|
|
||||||
|
|
||||||
envResult = environment.step(chosenAction);
|
envResult = environment.step(chosenAction);
|
||||||
State nextState = envResult.getState();
|
State nextState = envResult.getState();
|
||||||
@@ -86,12 +72,9 @@ public class MonteCarloControlEGreedy<A extends Enum> extends EpisodicLearning<A
|
|||||||
}
|
}
|
||||||
timestamp++;
|
timestamp++;
|
||||||
dispatchStepEnd();
|
dispatchStepEnd();
|
||||||
if(converged) return;
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if(currentEpisode % 2 == 1){
|
|
||||||
return;
|
|
||||||
}
|
|
||||||
|
|
||||||
// System.out.printf("Episode %d \t Reward: %f \n", currentEpisode, sumOfRewards);
|
// System.out.printf("Episode %d \t Reward: %f \n", currentEpisode, sumOfRewards);
|
||||||
HashMap<Pair<State, A>, List<Integer>> stateActionPairs = new LinkedHashMap<>();
|
HashMap<Pair<State, A>, List<Integer>> stateActionPairs = new LinkedHashMap<>();
|
||||||
|
|||||||
@@ -5,7 +5,14 @@ import core.algo.EpisodicLearning;
|
|||||||
import core.policy.EpsilonGreedyPolicy;
|
import core.policy.EpsilonGreedyPolicy;
|
||||||
import core.policy.GreedyPolicy;
|
import core.policy.GreedyPolicy;
|
||||||
import core.policy.Policy;
|
import core.policy.Policy;
|
||||||
|
import evironment.antGame.Reward;
|
||||||
|
import example.ContinuousAnt;
|
||||||
|
|
||||||
|
import java.io.File;
|
||||||
|
import java.io.IOException;
|
||||||
|
import java.nio.file.Files;
|
||||||
|
import java.nio.file.Path;
|
||||||
|
import java.nio.file.StandardOpenOption;
|
||||||
import java.util.Map;
|
import java.util.Map;
|
||||||
|
|
||||||
public class QLearningOffPolicyTDControl<A extends Enum> extends EpisodicLearning<A> {
|
public class QLearningOffPolicyTDControl<A extends Enum> extends EpisodicLearning<A> {
|
||||||
@@ -37,25 +44,57 @@ public class QLearningOffPolicyTDControl<A extends Enum> extends EpisodicLearnin
|
|||||||
|
|
||||||
|
|
||||||
sumOfRewards = 0;
|
sumOfRewards = 0;
|
||||||
|
int timestampTilFood = 0;
|
||||||
|
int foodCollected = 0;
|
||||||
|
int foodTimestampsTotal= 0;
|
||||||
while(envResult == null || !envResult.isDone()) {
|
while(envResult == null || !envResult.isDone()) {
|
||||||
actionValues = stateActionTable.getActionValues(state);
|
actionValues = stateActionTable.getActionValues(state);
|
||||||
A action;
|
A action = policy.chooseAction(actionValues);
|
||||||
if(currentEpisode % 2 == 0) {
|
|
||||||
action = greedyPolicy.chooseAction(actionValues);
|
|
||||||
} else {
|
|
||||||
action = policy.chooseAction(actionValues);
|
|
||||||
}
|
|
||||||
if(converged) return;
|
|
||||||
// Take a step
|
// Take a step
|
||||||
envResult = environment.step(action);
|
envResult = environment.step(action);
|
||||||
double reward = envResult.getReward();
|
double reward = envResult.getReward();
|
||||||
State nextState = envResult.getState();
|
State nextState = envResult.getState();
|
||||||
sumOfRewards += reward;
|
sumOfRewards += reward;
|
||||||
if(currentEpisode % 2 == 0) {
|
timestampTilFood++;
|
||||||
state = nextState;
|
|
||||||
dispatchStepEnd();
|
if(reward == Reward.FOOD_DROP_DOWN_SUCCESS) {
|
||||||
continue;
|
foodCollected++;
|
||||||
|
foodTimestampsTotal += timestampTilFood;
|
||||||
|
File file = new File(ContinuousAnt.FILE_NAME);
|
||||||
|
if(foodCollected % 1000 == 0) {
|
||||||
|
System.out.println(foodTimestampsTotal / 1000f + " " + timestampCurrentEpisode);
|
||||||
|
try {
|
||||||
|
Files.writeString(Path.of(file.getPath()), foodTimestampsTotal / 1000f + ",", StandardOpenOption.APPEND);
|
||||||
|
} catch (IOException e) {
|
||||||
|
e.printStackTrace();
|
||||||
|
}
|
||||||
|
foodTimestampsTotal = 0;
|
||||||
|
}
|
||||||
|
if(foodCollected == 1000){
|
||||||
|
((EpsilonGreedyPolicy<A>) this.policy).setEpsilon(0.15f);
|
||||||
|
}
|
||||||
|
if(foodCollected == 2000){
|
||||||
|
((EpsilonGreedyPolicy<A>) this.policy).setEpsilon(0.10f);
|
||||||
|
}
|
||||||
|
if(foodCollected == 3000){
|
||||||
|
((EpsilonGreedyPolicy<A>) this.policy).setEpsilon(0.05f);
|
||||||
|
}
|
||||||
|
if(foodCollected == 4000){
|
||||||
|
System.out.println("Reached 0 exploration");
|
||||||
|
((EpsilonGreedyPolicy<A>) this.policy).setEpsilon(0.00f);
|
||||||
|
}
|
||||||
|
if(foodCollected == 15000){
|
||||||
|
try {
|
||||||
|
Files.writeString(Path.of(file.getPath()), "\n", StandardOpenOption.APPEND);
|
||||||
|
} catch (IOException e) {
|
||||||
|
e.printStackTrace();
|
||||||
|
}
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
timestampTilFood = 0;
|
||||||
}
|
}
|
||||||
|
|
||||||
// Q Update
|
// Q Update
|
||||||
double currentQValue = stateActionTable.getActionValues(state).get(action);
|
double currentQValue = stateActionTable.getActionValues(state).get(action);
|
||||||
// maxQ(S', a);
|
// maxQ(S', a);
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ import java.util.Map;
|
|||||||
|
|
||||||
public class SARSA<A extends Enum> extends EpisodicLearning<A> {
|
public class SARSA<A extends Enum> extends EpisodicLearning<A> {
|
||||||
private float alpha;
|
private float alpha;
|
||||||
private Policy<A> greedyPolicy = new GreedyPolicy<>();
|
|
||||||
|
|
||||||
public SARSA(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, float learningRate, int delay) {
|
public SARSA(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, float learningRate, int delay) {
|
||||||
super(environment, actionSpace, discountFactor, delay);
|
super(environment, actionSpace, discountFactor, delay);
|
||||||
@@ -35,18 +34,13 @@ public class SARSA<A extends Enum> extends EpisodicLearning<A> {
|
|||||||
|
|
||||||
StepResultEnvironment envResult = null;
|
StepResultEnvironment envResult = null;
|
||||||
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
||||||
A action;
|
A action = policy.chooseAction(actionValues);
|
||||||
if(currentEpisode % 2 == 1){
|
|
||||||
action = greedyPolicy.chooseAction(actionValues);
|
|
||||||
}else{
|
|
||||||
action = policy.chooseAction(actionValues);
|
|
||||||
}
|
|
||||||
//A action = policy.chooseAction(actionValues);
|
//A action = policy.chooseAction(actionValues);
|
||||||
|
|
||||||
sumOfRewards = 0;
|
sumOfRewards = 0;
|
||||||
while(envResult == null || !envResult.isDone()) {
|
while(envResult == null || !envResult.isDone()) {
|
||||||
|
|
||||||
if(converged) return;
|
|
||||||
// Take a step
|
// Take a step
|
||||||
envResult = environment.step(action);
|
envResult = environment.step(action);
|
||||||
sumOfRewards += envResult.getReward();
|
sumOfRewards += envResult.getReward();
|
||||||
@@ -56,19 +50,8 @@ public class SARSA<A extends Enum> extends EpisodicLearning<A> {
|
|||||||
// Pick next action
|
// Pick next action
|
||||||
actionValues = stateActionTable.getActionValues(nextState);
|
actionValues = stateActionTable.getActionValues(nextState);
|
||||||
|
|
||||||
A nextAction;
|
A nextAction = policy.chooseAction(actionValues);
|
||||||
if(currentEpisode % 2 == 1){
|
|
||||||
nextAction = greedyPolicy.chooseAction(actionValues);
|
|
||||||
}else{
|
|
||||||
nextAction = policy.chooseAction(actionValues);
|
|
||||||
}
|
|
||||||
//A nextAction = policy.chooseAction(actionValues);
|
|
||||||
if(currentEpisode % 2 == 1){
|
|
||||||
state = nextState;
|
|
||||||
action = nextAction;
|
|
||||||
dispatchStepEnd();
|
|
||||||
continue;
|
|
||||||
}
|
|
||||||
// td update
|
// td update
|
||||||
// target = reward + gamma * Q(nextState, nextAction)
|
// target = reward + gamma * Q(nextState, nextAction)
|
||||||
double currentQValue = stateActionTable.getActionValues(state).get(action);
|
double currentQValue = stateActionTable.getActionValues(state).get(action);
|
||||||
|
|||||||
@@ -54,7 +54,7 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
|||||||
|
|
||||||
protected StepCalculation processStep(AntAction action) {
|
protected StepCalculation processStep(AntAction action) {
|
||||||
StepCalculation sc = new StepCalculation();
|
StepCalculation sc = new StepCalculation();
|
||||||
sc.reward = -1;
|
sc.reward = Reward.DEFAULT_REWARD;
|
||||||
sc.info = "";
|
sc.info = "";
|
||||||
sc.done = false;
|
sc.done = false;
|
||||||
Cell currentCell = grid.getCell(myAnt.getPos());
|
Cell currentCell = grid.getCell(myAnt.getPos());
|
||||||
@@ -84,10 +84,10 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
|||||||
case PICK_UP:
|
case PICK_UP:
|
||||||
if(myAnt.hasFood()) {
|
if(myAnt.hasFood()) {
|
||||||
// Ant tries to pick up food but can only hold one piece
|
// Ant tries to pick up food but can only hold one piece
|
||||||
sc.reward += Reward.FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY;
|
sc.reward = Reward.FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY;
|
||||||
} else if(currentCell.getFood() == 0) {
|
} else if(currentCell.getFood() == 0) {
|
||||||
// Ant tries to pick up food on cell that has no food on it
|
// Ant tries to pick up food on cell that has no food on it
|
||||||
sc.reward += Reward.FOOD_PICK_UP_FAIL_NO_FOOD;
|
sc.reward = Reward.FOOD_PICK_UP_FAIL_NO_FOOD;
|
||||||
} else if(currentCell.getFood() > 0) {
|
} else if(currentCell.getFood() > 0) {
|
||||||
// Ant successfully picks up food
|
// Ant successfully picks up food
|
||||||
currentCell.setFood(currentCell.getFood() - 1);
|
currentCell.setFood(currentCell.getFood() - 1);
|
||||||
@@ -98,13 +98,13 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
|||||||
case DROP_DOWN:
|
case DROP_DOWN:
|
||||||
if(!myAnt.hasFood()) {
|
if(!myAnt.hasFood()) {
|
||||||
// Ant had no food to drop
|
// Ant had no food to drop
|
||||||
sc.reward += Reward.FOOD_DROP_DOWN_FAIL_NO_FOOD;
|
sc.reward = Reward.FOOD_DROP_DOWN_FAIL_NO_FOOD;
|
||||||
} else {
|
} else {
|
||||||
myAnt.setHasFood(false);
|
myAnt.setHasFood(false);
|
||||||
// negative reward if the agent drops food on any other field
|
// negative reward if the agent drops food on any other field
|
||||||
// than the starting point
|
// than the starting point
|
||||||
if(currentCell.getType() != CellType.START) {
|
if(currentCell.getType() != CellType.START) {
|
||||||
sc.reward += Reward.FOOD_DROP_DOWN_FAIL_NOT_START;
|
sc.reward = Reward.FOOD_DROP_DOWN_FAIL_NOT_START;
|
||||||
// Drop food onto the ground
|
// Drop food onto the ground
|
||||||
currentCell.setFood(currentCell.getFood() + 1);
|
currentCell.setFood(currentCell.getFood() + 1);
|
||||||
} else {
|
} else {
|
||||||
@@ -122,10 +122,10 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
|||||||
if(!sc.stayOnCell) {
|
if(!sc.stayOnCell) {
|
||||||
if(!isInGrid(sc.potentialNextPos)) {
|
if(!isInGrid(sc.potentialNextPos)) {
|
||||||
sc.stayOnCell = true;
|
sc.stayOnCell = true;
|
||||||
sc.reward += Reward.RAN_INTO_WALL;
|
sc.reward = Reward.RAN_INTO_WALL;
|
||||||
} else if(hitObstacle(sc.potentialNextPos)) {
|
} else if(hitObstacle(sc.potentialNextPos)) {
|
||||||
sc.stayOnCell = true;
|
sc.stayOnCell = true;
|
||||||
sc.reward += Reward.RAN_INTO_OBSTACLE;
|
sc.reward = Reward.RAN_INTO_OBSTACLE;
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -60,8 +60,8 @@ public class Grid {
|
|||||||
Point potFood = new Point(0, 0);
|
Point potFood = new Point(0, 0);
|
||||||
CellType potFieldType;
|
CellType potFieldType;
|
||||||
while(!foodSpawned) {
|
while(!foodSpawned) {
|
||||||
potFood.x = RNG.getEnvRandom().nextInt(width);
|
potFood.x = RNG.getRandomEnv().nextInt(width);
|
||||||
potFood.y = RNG.getEnvRandom().nextInt(height);
|
potFood.y = RNG.getRandomEnv().nextInt(height);
|
||||||
potFieldType = grid[potFood.x][potFood.y].getType();
|
potFieldType = grid[potFood.x][potFood.y].getType();
|
||||||
if(potFieldType != CellType.START && grid[potFood.x][potFood.y].getFood() == 0 && potFieldType != CellType.OBSTACLE) {
|
if(potFieldType != CellType.START && grid[potFood.x][potFood.y].getFood() == 0 && potFieldType != CellType.OBSTACLE) {
|
||||||
grid[potFood.x][potFood.y].setFood(1);
|
grid[potFood.x][potFood.y].setFood(1);
|
||||||
|
|||||||
@@ -1,16 +1,17 @@
|
|||||||
package evironment.antGame;
|
package evironment.antGame;
|
||||||
|
|
||||||
public class Reward {
|
public class Reward {
|
||||||
public static final double FOOD_PICK_UP_SUCCESS = 1;
|
public static final double DEFAULT_REWARD = -1;
|
||||||
public static final double FOOD_PICK_UP_FAIL_NO_FOOD = -1;
|
public static final double FOOD_PICK_UP_SUCCESS = 0;
|
||||||
public static final double FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY = -1;
|
public static final double FOOD_PICK_UP_FAIL_NO_FOOD = -2;
|
||||||
|
public static final double FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY = -2;
|
||||||
|
|
||||||
public static final double FOOD_DROP_DOWN_FAIL_NO_FOOD = -1;
|
public static final double FOOD_DROP_DOWN_FAIL_NO_FOOD = -2;
|
||||||
public static final double FOOD_DROP_DOWN_FAIL_NOT_START = -1;
|
public static final double FOOD_DROP_DOWN_FAIL_NOT_START = -2;
|
||||||
public static final double FOOD_DROP_DOWN_SUCCESS = 40;
|
public static final double FOOD_DROP_DOWN_SUCCESS = 1;
|
||||||
|
|
||||||
public static final double UNKNOWN_FIELD_EXPLORED = 0;
|
public static final double UNKNOWN_FIELD_EXPLORED = 0;
|
||||||
|
|
||||||
public static final double RAN_INTO_WALL = -1;
|
public static final double RAN_INTO_WALL = -2;
|
||||||
public static final double RAN_INTO_OBSTACLE = -1;
|
public static final double RAN_INTO_OBSTACLE = -2;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -29,6 +29,6 @@ public class CardDeck {
|
|||||||
nextInt(int bound) returns random int value from (inclusive) 0
|
nextInt(int bound) returns random int value from (inclusive) 0
|
||||||
and EXCLUSIVE! bound
|
and EXCLUSIVE! bound
|
||||||
*/
|
*/
|
||||||
return cards.get(RNG.getEnvRandom().nextInt(cards.size()));
|
return cards.get(RNG.getRandomEnv().nextInt(cards.size()));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ public class DinoWorldAdvanced extends DinoWorld{
|
|||||||
protected void spawnNewObstacle() {
|
protected void spawnNewObstacle() {
|
||||||
int dx;
|
int dx;
|
||||||
int xSpawn;
|
int xSpawn;
|
||||||
double ran = RNG.getEnvRandom().nextDouble();
|
double ran = RNG.getRandomEnv().nextDouble();
|
||||||
if(ran < 0.25){
|
if(ran < 0.25){
|
||||||
dx = -(int) (0.35 * Config.OBSTACLE_SPEED);
|
dx = -(int) (0.35 * Config.OBSTACLE_SPEED);
|
||||||
}else if(ran < 0.5){
|
}else if(ran < 0.5){
|
||||||
@@ -41,7 +41,7 @@ public class DinoWorldAdvanced extends DinoWorld{
|
|||||||
} else{
|
} else{
|
||||||
dx = -(int) (3.5 * Config.OBSTACLE_SPEED);
|
dx = -(int) (3.5 * Config.OBSTACLE_SPEED);
|
||||||
}
|
}
|
||||||
double ran2 = RNG.getEnvRandom().nextDouble();
|
double ran2 = RNG.getRandomEnv().nextDouble();
|
||||||
if(ran2 < 0.25) {
|
if(ran2 < 0.25) {
|
||||||
// randomly spawning more right outside of the screen
|
// randomly spawning more right outside of the screen
|
||||||
xSpawn = Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE;
|
xSpawn = Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE;
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ import evironment.blackjack.PlayerAction;
|
|||||||
|
|
||||||
public class BlackJack {
|
public class BlackJack {
|
||||||
public static void main(String[] args) {
|
public static void main(String[] args) {
|
||||||
RNG.setSeed(55, true);
|
RNG.setSeed(55);
|
||||||
|
|
||||||
RLController<PlayerAction> rl = new RLControllerGUI<>(
|
RLController<PlayerAction> rl = new RLControllerGUI<>(
|
||||||
new BlackJackTable(),
|
new BlackJackTable(),
|
||||||
|
|||||||
@@ -0,0 +1,37 @@
|
|||||||
|
package example;
|
||||||
|
|
||||||
|
import core.RNG;
|
||||||
|
import core.algo.Method;
|
||||||
|
import core.controller.RLController;
|
||||||
|
import core.controller.RLControllerGUI;
|
||||||
|
import evironment.antGame.AntAction;
|
||||||
|
import evironment.antGame.AntWorldContinuous;
|
||||||
|
|
||||||
|
import java.io.File;
|
||||||
|
import java.io.IOException;
|
||||||
|
|
||||||
|
public class ContinuousAnt {
|
||||||
|
public static final String FILE_NAME = "converge.txt";
|
||||||
|
|
||||||
|
public static void main(String[] args) {
|
||||||
|
File file = new File(FILE_NAME);
|
||||||
|
try {
|
||||||
|
file.createNewFile();
|
||||||
|
} catch (IOException e) {
|
||||||
|
e.printStackTrace();
|
||||||
|
}
|
||||||
|
RNG.setSeed(13, true);
|
||||||
|
RLController<AntAction> rl = new RLControllerGUI<>(
|
||||||
|
new AntWorldContinuous(8, 8),
|
||||||
|
Method.Q_LEARNING_OFF_POLICY_CONTROL,
|
||||||
|
AntAction.values());
|
||||||
|
rl.setDelay(20);
|
||||||
|
rl.setNrOfEpisodes(1);
|
||||||
|
// 0.05, 0.1, 0.3, 0.5, 0.7, 0.9, 0.95, 0.99
|
||||||
|
rl.setDiscountFactor(0.05f);
|
||||||
|
// 0.1, 0.3, 0.5, 0.7 0.9
|
||||||
|
rl.setLearningRate(0.9f);
|
||||||
|
rl.setEpsilon(0.2f);
|
||||||
|
rl.start();
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,52 +0,0 @@
|
|||||||
package example;
|
|
||||||
|
|
||||||
import core.RNG;
|
|
||||||
import core.algo.Method;
|
|
||||||
import core.controller.RLController;
|
|
||||||
import evironment.jumpingDino.DinoAction;
|
|
||||||
import evironment.jumpingDino.DinoWorldAdvanced;
|
|
||||||
|
|
||||||
import java.io.File;
|
|
||||||
import java.io.IOException;
|
|
||||||
import java.nio.file.Files;
|
|
||||||
import java.nio.file.Path;
|
|
||||||
import java.nio.file.StandardOpenOption;
|
|
||||||
|
|
||||||
public class DinoSampling {
|
|
||||||
public static final String FILE_NAME = "advancedEveryVisit.txt";
|
|
||||||
public static void main(String[] args) {
|
|
||||||
File file = new File(FILE_NAME);
|
|
||||||
try {
|
|
||||||
file.createNewFile();
|
|
||||||
} catch (IOException e) {
|
|
||||||
e.printStackTrace();
|
|
||||||
}
|
|
||||||
for(float f = 0.05f; f <= 1.003; f += 0.05f) {
|
|
||||||
try {
|
|
||||||
Files.writeString(Path.of(file.getPath()), f + ",", StandardOpenOption.APPEND);
|
|
||||||
} catch (IOException e) {
|
|
||||||
e.printStackTrace();
|
|
||||||
}
|
|
||||||
for(int i = 1; i <= 100; i++) {
|
|
||||||
System.out.println("seed: " + i * 13);
|
|
||||||
RNG.setSeed(i * 13, true);
|
|
||||||
|
|
||||||
RLController<DinoAction> rl = new RLController<>(
|
|
||||||
new DinoWorldAdvanced(),
|
|
||||||
Method.MC_CONTROL_EVERY_VISIT,
|
|
||||||
DinoAction.values());
|
|
||||||
rl.setDelay(0);
|
|
||||||
rl.setDiscountFactor(1f);
|
|
||||||
rl.setEpsilon(f);
|
|
||||||
rl.setLearningRate(1f);
|
|
||||||
rl.setNrOfEpisodes(400000);
|
|
||||||
rl.start();
|
|
||||||
}
|
|
||||||
try {
|
|
||||||
Files.writeString(Path.of(file.getPath()), "\n", StandardOpenOption.APPEND);
|
|
||||||
} catch (IOException e) {
|
|
||||||
e.printStackTrace();
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -4,11 +4,12 @@ import core.RNG;
|
|||||||
import core.algo.Method;
|
import core.algo.Method;
|
||||||
import core.controller.RLController;
|
import core.controller.RLController;
|
||||||
import evironment.jumpingDino.DinoAction;
|
import evironment.jumpingDino.DinoAction;
|
||||||
|
import evironment.jumpingDino.DinoWorld;
|
||||||
import evironment.jumpingDino.DinoWorldAdvanced;
|
import evironment.jumpingDino.DinoWorldAdvanced;
|
||||||
|
|
||||||
public class JumpingDino {
|
public class JumpingDino {
|
||||||
public static void main(String[] args) {
|
public static void main(String[] args) {
|
||||||
RNG.setSeed(29, true);
|
RNG.setSeed(29);
|
||||||
|
|
||||||
RLController<DinoAction> rl = new RLController<>(
|
RLController<DinoAction> rl = new RLController<>(
|
||||||
new DinoWorldAdvanced(),
|
new DinoWorldAdvanced(),
|
||||||
@@ -16,9 +17,9 @@ public class JumpingDino {
|
|||||||
DinoAction.values());
|
DinoAction.values());
|
||||||
|
|
||||||
rl.setDelay(0);
|
rl.setDelay(0);
|
||||||
rl.setDiscountFactor(1f);
|
rl.setDiscountFactor(9f);
|
||||||
rl.setEpsilon(0.05f);
|
rl.setEpsilon(0.05f);
|
||||||
rl.setLearningRate(1f);
|
rl.setLearningRate(0.8f);
|
||||||
rl.setNrOfEpisodes(100000);
|
rl.setNrOfEpisodes(100000);
|
||||||
rl.start();
|
rl.start();
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ import evironment.antGame.AntWorld;
|
|||||||
|
|
||||||
public class RunningAnt {
|
public class RunningAnt {
|
||||||
public static void main(String[] args) {
|
public static void main(String[] args) {
|
||||||
RNG.setSeed(56, true);
|
RNG.setSeed(56);
|
||||||
|
|
||||||
RLController<AntAction> rl = new RLControllerGUI<>(
|
RLController<AntAction> rl = new RLControllerGUI<>(
|
||||||
new AntWorld(8, 8),
|
new AntWorld(8, 8),
|
||||||
@@ -19,6 +19,7 @@ public class RunningAnt {
|
|||||||
rl.setDelay(200);
|
rl.setDelay(200);
|
||||||
rl.setNrOfEpisodes(10000);
|
rl.setNrOfEpisodes(10000);
|
||||||
rl.setDiscountFactor(0.9f);
|
rl.setDiscountFactor(0.9f);
|
||||||
|
rl.setLearningRate(0.9f);
|
||||||
rl.setEpsilon(0.15f);
|
rl.setEpsilon(0.15f);
|
||||||
rl.start();
|
rl.start();
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user