Merge branch 'antWorldRewardAnalysis'
# Conflicts: # src/main/java/core/algo/EpisodicLearning.java # src/main/java/core/controller/RLController.java # src/main/java/evironment/jumpingDino/DinoWorld.java # src/main/java/evironment/jumpingDino/DinoWorldAdvanced.java # src/main/java/example/JumpingDino.java
This commit is contained in:
@@ -1,9 +1,10 @@
|
||||
package core;
|
||||
|
||||
import java.security.SecureRandom;
|
||||
import java.util.Random;
|
||||
|
||||
/**
|
||||
* ! SecureRandom not working properly on windows/different JDKs,
|
||||
* using Random again !
|
||||
*
|
||||
* To ensure deterministic behaviour of repeating program executions,
|
||||
* this class is used for all random number generation methods.
|
||||
* Do not use Math.random()!
|
||||
@@ -13,19 +14,33 @@ import java.util.Random;
|
||||
* execution)
|
||||
*/
|
||||
public class RNG {
|
||||
private static SecureRandom rng;
|
||||
private static Random rng;
|
||||
private static Random rngEnv;
|
||||
private static int seed = 123;
|
||||
private static int envSeed = 13;
|
||||
static {
|
||||
rng = new SecureRandom();
|
||||
rng.setSeed(seed);
|
||||
rng = new Random();
|
||||
rngEnv = new Random();
|
||||
setSeed(seed, true);
|
||||
}
|
||||
|
||||
public static Random getRandom() {
|
||||
return rng;
|
||||
}
|
||||
public static Random getRandomEnv() {
|
||||
return rngEnv;
|
||||
}
|
||||
|
||||
public static void setSeed(int seed){
|
||||
public static void setSeed(int seed, boolean setEnvRandom) {
|
||||
RNG.seed = seed;
|
||||
rng.setSeed(seed);
|
||||
if(setEnvRandom) {
|
||||
rngEnv.setSeed(13);
|
||||
}
|
||||
}
|
||||
|
||||
public static void setSeed(int seed) {
|
||||
setSeed(seed, true);
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ import core.Environment;
|
||||
import core.LearningConfig;
|
||||
import core.StepResult;
|
||||
import core.listener.LearningListener;
|
||||
import core.policy.EpsilonGreedyPolicy;
|
||||
import lombok.Getter;
|
||||
import lombok.Setter;
|
||||
|
||||
@@ -73,7 +72,7 @@ public abstract class EpisodicLearning<A extends Enum> extends Learning<A> imple
|
||||
}
|
||||
}
|
||||
|
||||
protected void dispatchEpisodeStart(){
|
||||
private void dispatchEpisodeStart(){
|
||||
++currentEpisode;
|
||||
episodesToLearn.decrementAndGet();
|
||||
for(LearningListener l: learningListeners){
|
||||
@@ -85,13 +84,20 @@ public abstract class EpisodicLearning<A extends Enum> extends Learning<A> imple
|
||||
protected void dispatchStepEnd() {
|
||||
super.dispatchStepEnd();
|
||||
timestamp++;
|
||||
timestampCurrentEpisode++;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void learn(){
|
||||
learn(LearningConfig.DEFAULT_NR_OF_EPISODES);
|
||||
}
|
||||
|
||||
private void startLearning(){
|
||||
dispatchLearningStart();
|
||||
System.out.println(episodesToLearn.get());
|
||||
while(episodesToLearn.get() > 0){
|
||||
|
||||
dispatchEpisodeStart();
|
||||
timestampCurrentEpisode = 0;
|
||||
nextEpisode();
|
||||
dispatchEpisodeEnd();
|
||||
}
|
||||
@@ -104,7 +110,6 @@ public abstract class EpisodicLearning<A extends Enum> extends Learning<A> imple
|
||||
public void learnMoreEpisodes(int nrOfEpisodes){
|
||||
episodesToLearn.addAndGet(nrOfEpisodes);
|
||||
}
|
||||
|
||||
/**
|
||||
* Stopping the while loop by setting episodesToLearn to 0.
|
||||
* The current episode can not be interrupted, so the sleep delay
|
||||
@@ -127,14 +132,8 @@ public abstract class EpisodicLearning<A extends Enum> extends Learning<A> imple
|
||||
delay = prevDelay;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void learn(){
|
||||
learn(LearningConfig.DEFAULT_NR_OF_EPISODES);
|
||||
}
|
||||
|
||||
public synchronized void learn(int nrOfEpisodes){
|
||||
boolean isLearning = episodesToLearn.getAndAdd(nrOfEpisodes) != 0;
|
||||
System.out.println(isLearning);
|
||||
if(!isLearning)
|
||||
startLearning();
|
||||
}
|
||||
|
||||
@@ -16,8 +16,6 @@ import java.util.HashSet;
|
||||
import java.util.List;
|
||||
import java.util.Set;
|
||||
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 {
|
||||
oos.writeObject(rewardHistory);
|
||||
oos.writeObject(stateActionTable);
|
||||
// oos.writeObject(stateActionTable);
|
||||
}
|
||||
|
||||
public void load(ObjectInputStream ois) throws IOException, ClassNotFoundException {
|
||||
|
||||
@@ -5,5 +5,5 @@ package core.algo;
|
||||
* which RL-algorithm should be used.
|
||||
*/
|
||||
public enum Method {
|
||||
MC_CONTROL_FIRST_VISIT, SARSA_EPISODIC, Q_LEARNING_OFF_POLICY_CONTROL
|
||||
MC_CONTROL_FIRST_VISIT, MC_CONTROL_EVERY_VISIT, SARSA_ON_POLICY_CONTROL, Q_LEARNING_OFF_POLICY_CONTROL
|
||||
}
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
package core.algo.mc;
|
||||
|
||||
import core.*;
|
||||
import core.algo.EpisodicLearning;
|
||||
import core.policy.EpsilonGreedyPolicy;
|
||||
import org.apache.commons.lang3.tuple.ImmutablePair;
|
||||
import org.apache.commons.lang3.tuple.Pair;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.ObjectInputStream;
|
||||
import java.io.ObjectOutputStream;
|
||||
import java.util.*;
|
||||
|
||||
/**
|
||||
* Includes both! variants of Monte-Carlo methods
|
||||
* Default method is First-Visit.
|
||||
* Change to Every-Visit by setting flag "useEveryVisit" in the constructor to true.
|
||||
* @param <A>
|
||||
*/
|
||||
public class MonteCarloControlEGreedy<A extends Enum> extends EpisodicLearning<A> {
|
||||
|
||||
private Map<Pair<State, A>, Double> returnSum;
|
||||
private Map<Pair<State, A>, Integer> returnCount;
|
||||
private boolean isEveryVisit;
|
||||
|
||||
|
||||
public MonteCarloControlEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, int delay, boolean useEveryVisit) {
|
||||
super(environment, actionSpace, discountFactor, delay);
|
||||
isEveryVisit = useEveryVisit;
|
||||
this.policy = new EpsilonGreedyPolicy<>(epsilon);
|
||||
this.stateActionTable = new DeterministicStateActionTable<>(this.actionSpace);
|
||||
returnSum = new HashMap<>();
|
||||
returnCount = new HashMap<>();
|
||||
}
|
||||
|
||||
public MonteCarloControlEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, int delay) {
|
||||
this(environment, actionSpace, discountFactor, epsilon, delay, false);
|
||||
}
|
||||
|
||||
public MonteCarloControlEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, int delay) {
|
||||
this(environment, actionSpace, LearningConfig.DEFAULT_DISCOUNT_FACTOR, LearningConfig.DEFAULT_EPSILON, delay);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void nextEpisode() {
|
||||
episode = new ArrayList<>();
|
||||
State state = environment.reset();
|
||||
try {
|
||||
Thread.sleep(delay);
|
||||
} catch (InterruptedException e) {
|
||||
e.printStackTrace();
|
||||
}
|
||||
sumOfRewards = 0;
|
||||
StepResultEnvironment envResult = null;
|
||||
|
||||
while(envResult == null || !envResult.isDone()) {
|
||||
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
||||
A chosenAction = policy.chooseAction(actionValues);
|
||||
|
||||
envResult = environment.step(chosenAction);
|
||||
State nextState = envResult.getState();
|
||||
sumOfRewards += envResult.getReward();
|
||||
rewardCheckSum += envResult.getReward();
|
||||
episode.add(new StepResult<>(state, chosenAction, envResult.getReward()));
|
||||
|
||||
state = nextState;
|
||||
|
||||
try {
|
||||
Thread.sleep(delay);
|
||||
} catch (InterruptedException e) {
|
||||
e.printStackTrace();
|
||||
}
|
||||
timestamp++;
|
||||
dispatchStepEnd();
|
||||
}
|
||||
|
||||
|
||||
|
||||
// System.out.printf("Episode %d \t Reward: %f \n", currentEpisode, sumOfRewards);
|
||||
HashMap<Pair<State, A>, List<Integer>> stateActionPairs = new LinkedHashMap<>();
|
||||
|
||||
int firstOccurrenceIndex = 0;
|
||||
for(StepResult<A> sr : episode) {
|
||||
Pair<State, A> pair = new ImmutablePair<>(sr.getState(), sr.getAction());
|
||||
if(!stateActionPairs.containsKey(pair)) {
|
||||
List<Integer> l = new ArrayList<>();
|
||||
l.add(firstOccurrenceIndex);
|
||||
stateActionPairs.put(pair, l);
|
||||
}
|
||||
|
||||
/*
|
||||
This is the only difference between First-Visit and Every-Visit.
|
||||
When First-Visit is selected, only the first index of the occurrence is put into the list.
|
||||
When Every-Visit is selected, every following occurrence is saved
|
||||
into the list as well.
|
||||
*/
|
||||
else if(isEveryVisit) {
|
||||
stateActionPairs.get(pair).add(firstOccurrenceIndex);
|
||||
}
|
||||
++firstOccurrenceIndex;
|
||||
}
|
||||
//System.out.println("stateActionPairs " + stateActionPairs.size());
|
||||
for(Map.Entry<Pair<State, A>, List<Integer>> entry : stateActionPairs.entrySet()) {
|
||||
Pair<State, A> stateActionPair = entry.getKey();
|
||||
List<Integer> firstOccurrences = entry.getValue();
|
||||
for(Integer firstOccurrencesIdx : firstOccurrences) {
|
||||
double G = 0;
|
||||
for(int l = firstOccurrencesIdx; l < episode.size(); ++l) {
|
||||
G += episode.get(l).getReward() * (Math.pow(discountFactor, l - firstOccurrencesIdx));
|
||||
}
|
||||
// slick trick to add G to the entry.
|
||||
// if the key does not exists, it will create a new entry with G as default value
|
||||
returnSum.merge(stateActionPair, G, Double::sum);
|
||||
returnCount.merge(stateActionPair, 1, Integer::sum);
|
||||
stateActionTable.setValue(stateActionPair.getKey(), stateActionPair.getValue(), returnSum.get(stateActionPair) / returnCount.get(stateActionPair));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@Override
|
||||
public void save(ObjectOutputStream oos) throws IOException {
|
||||
super.save(oos);
|
||||
oos.writeObject(returnSum);
|
||||
oos.writeObject(returnCount);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void load(ObjectInputStream ois) throws IOException, ClassNotFoundException {
|
||||
super.load(ois);
|
||||
returnSum = (Map<Pair<State, A>, Double>) ois.readObject();
|
||||
returnCount = (Map<Pair<State, A>, Integer>) ois.readObject();
|
||||
}
|
||||
}
|
||||
@@ -1,127 +0,0 @@
|
||||
package core.algo.mc;
|
||||
|
||||
import core.*;
|
||||
import core.algo.EpisodicLearning;
|
||||
import core.policy.EpsilonGreedyPolicy;
|
||||
import org.apache.commons.lang3.tuple.ImmutablePair;
|
||||
import org.apache.commons.lang3.tuple.Pair;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.ObjectInputStream;
|
||||
import java.io.ObjectOutputStream;
|
||||
import java.util.*;
|
||||
|
||||
/**
|
||||
* TODO: Major problem:
|
||||
* StateActionPairs are only unique accounting for their position in the episode.
|
||||
* For example:
|
||||
* <p>
|
||||
* startingState -> MOVE_LEFT : very first state action in the episode i = 1
|
||||
* image the agent does not collect the food and does not drop it onto start, the agent will receive
|
||||
* -1 for every timestamp hence (startingState -> MOVE_LEFT) will get a value of -10;
|
||||
* <p>
|
||||
* BUT image moving left from the starting position will have no impact on the state because
|
||||
* the agent ran into a wall. The known world stays the same.
|
||||
* Taking an action after that will have the exact same state but a different action
|
||||
* making the value of this stateActionPair -9 because the stateAction pair took place on the second
|
||||
* timestamp, summing up all remaining rewards will be -9...
|
||||
* <p>
|
||||
* How to encounter this problem?
|
||||
*
|
||||
* @param <A>
|
||||
*/
|
||||
public class MonteCarloControlFirstVisitEGreedy<A extends Enum> extends EpisodicLearning<A> {
|
||||
|
||||
private Map<Pair<State, A>, Double> returnSum;
|
||||
private Map<Pair<State, A>, Integer> returnCount;
|
||||
|
||||
public MonteCarloControlFirstVisitEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, int delay) {
|
||||
super(environment, actionSpace, discountFactor, delay);
|
||||
this.policy = new EpsilonGreedyPolicy<>(epsilon);
|
||||
this.stateActionTable = new DeterministicStateActionTable<>(this.actionSpace);
|
||||
returnSum = new HashMap<>();
|
||||
returnCount = new HashMap<>();
|
||||
}
|
||||
|
||||
public MonteCarloControlFirstVisitEGreedy(Environment<A> environment, DiscreteActionSpace<A> actionSpace, int delay) {
|
||||
this(environment, actionSpace, LearningConfig.DEFAULT_DISCOUNT_FACTOR, LearningConfig.DEFAULT_EPSILON, delay);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void nextEpisode() {
|
||||
episode = new ArrayList<>();
|
||||
State state = environment.reset();
|
||||
try {
|
||||
Thread.sleep(delay);
|
||||
} catch (InterruptedException e) {
|
||||
e.printStackTrace();
|
||||
}
|
||||
sumOfRewards = 0;
|
||||
StepResultEnvironment envResult = null;
|
||||
//TODO extract to learning
|
||||
int timestamp = 0;
|
||||
while(envResult == null || !envResult.isDone()) {
|
||||
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
||||
A chosenAction = policy.chooseAction(actionValues);
|
||||
checkSum += chosenAction.ordinal();
|
||||
envResult = environment.step(chosenAction);
|
||||
State nextState = envResult.getState();
|
||||
sumOfRewards += envResult.getReward();
|
||||
rewardCheckSum += envResult.getReward();
|
||||
episode.add(new StepResult<>(state, chosenAction, envResult.getReward()));
|
||||
|
||||
state = nextState;
|
||||
|
||||
try {
|
||||
Thread.sleep(delay);
|
||||
} catch (InterruptedException e) {
|
||||
e.printStackTrace();
|
||||
}
|
||||
timestamp++;
|
||||
dispatchStepEnd();
|
||||
}
|
||||
|
||||
// System.out.printf("Episode %d \t Reward: %f \n", currentEpisode, sumOfRewards);
|
||||
Set<Pair<State, A>> stateActionPairs = new LinkedHashSet<>();
|
||||
|
||||
for(StepResult<A> sr : episode) {
|
||||
stateActionPairs.add(new ImmutablePair<>(sr.getState(), sr.getAction()));
|
||||
}
|
||||
|
||||
//System.out.println("stateActionPairs " + stateActionPairs.size());
|
||||
for(Pair<State, A> stateActionPair : stateActionPairs) {
|
||||
int firstOccurenceIndex = 0;
|
||||
// find first occurance of state action pair
|
||||
for(StepResult<A> sr : episode) {
|
||||
if(stateActionPair.getKey().equals(sr.getState()) && stateActionPair.getValue().equals(sr.getAction())) {
|
||||
break;
|
||||
}
|
||||
firstOccurenceIndex++;
|
||||
}
|
||||
|
||||
double G = 0;
|
||||
for(int l = firstOccurenceIndex; l < episode.size(); ++l) {
|
||||
G += episode.get(l).getReward() * (Math.pow(discountFactor, l - firstOccurenceIndex));
|
||||
}
|
||||
// slick trick to add G to the entry.
|
||||
// if the key does not exists, it will create a new entry with G as default value
|
||||
returnSum.merge(stateActionPair, G, Double::sum);
|
||||
returnCount.merge(stateActionPair, 1, Integer::sum);
|
||||
stateActionTable.setValue(stateActionPair.getKey(), stateActionPair.getValue(), returnSum.get(stateActionPair) / returnCount.get(stateActionPair));
|
||||
}
|
||||
}
|
||||
|
||||
@Override
|
||||
public void save(ObjectOutputStream oos) throws IOException {
|
||||
super.save(oos);
|
||||
oos.writeObject(returnSum);
|
||||
oos.writeObject(returnCount);
|
||||
}
|
||||
|
||||
@Override
|
||||
public void load(ObjectInputStream ois) throws IOException, ClassNotFoundException {
|
||||
super.load(ois);
|
||||
returnSum = (Map<Pair<State, A>, Double>) ois.readObject();
|
||||
returnCount = (Map<Pair<State, A>, Integer>) ois.readObject();
|
||||
}
|
||||
}
|
||||
@@ -5,7 +5,14 @@ import core.algo.EpisodicLearning;
|
||||
import core.policy.EpsilonGreedyPolicy;
|
||||
import core.policy.GreedyPolicy;
|
||||
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;
|
||||
|
||||
public class QLearningOffPolicyTDControl<A extends Enum> extends EpisodicLearning<A> {
|
||||
@@ -37,6 +44,9 @@ public class QLearningOffPolicyTDControl<A extends Enum> extends EpisodicLearnin
|
||||
|
||||
|
||||
sumOfRewards = 0;
|
||||
int timestampTilFood = 0;
|
||||
int foodCollected = 0;
|
||||
int foodTimestampsTotal= 0;
|
||||
while(envResult == null || !envResult.isDone()) {
|
||||
actionValues = stateActionTable.getActionValues(state);
|
||||
A action = policy.chooseAction(actionValues);
|
||||
@@ -46,6 +56,44 @@ public class QLearningOffPolicyTDControl<A extends Enum> extends EpisodicLearnin
|
||||
double reward = envResult.getReward();
|
||||
State nextState = envResult.getState();
|
||||
sumOfRewards += reward;
|
||||
timestampTilFood++;
|
||||
|
||||
if(reward == Reward.FOOD_DROP_DOWN_SUCCESS) {
|
||||
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
|
||||
double currentQValue = stateActionTable.getActionValues(state).get(action);
|
||||
|
||||
@@ -3,6 +3,8 @@ package core.algo.td;
|
||||
import core.*;
|
||||
import core.algo.EpisodicLearning;
|
||||
import core.policy.EpsilonGreedyPolicy;
|
||||
import core.policy.GreedyPolicy;
|
||||
import core.policy.Policy;
|
||||
|
||||
import java.util.Map;
|
||||
|
||||
@@ -32,10 +34,13 @@ public class SARSA<A extends Enum> extends EpisodicLearning<A> {
|
||||
|
||||
StepResultEnvironment envResult = null;
|
||||
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
|
||||
A action = policy.chooseAction(actionValues);
|
||||
A action = policy.chooseAction(actionValues);
|
||||
|
||||
//A action = policy.chooseAction(actionValues);
|
||||
|
||||
sumOfRewards = 0;
|
||||
while(envResult == null || !envResult.isDone()) {
|
||||
|
||||
// Take a step
|
||||
envResult = environment.step(action);
|
||||
sumOfRewards += envResult.getReward();
|
||||
@@ -44,7 +49,8 @@ public class SARSA<A extends Enum> extends EpisodicLearning<A> {
|
||||
|
||||
// Pick next action
|
||||
actionValues = stateActionTable.getActionValues(nextState);
|
||||
A nextAction = policy.chooseAction(actionValues);
|
||||
|
||||
A nextAction = policy.chooseAction(actionValues);
|
||||
|
||||
// td update
|
||||
// target = reward + gamma * Q(nextState, nextAction)
|
||||
|
||||
@@ -7,7 +7,7 @@ import core.ListDiscreteActionSpace;
|
||||
import core.algo.EpisodicLearning;
|
||||
import core.algo.Learning;
|
||||
import core.algo.Method;
|
||||
import core.algo.mc.MonteCarloControlFirstVisitEGreedy;
|
||||
import core.algo.mc.MonteCarloControlEGreedy;
|
||||
import core.algo.td.QLearningOffPolicyTDControl;
|
||||
import core.algo.td.SARSA;
|
||||
import core.listener.LearningListener;
|
||||
@@ -49,9 +49,13 @@ public class RLController<A extends Enum> implements LearningListener {
|
||||
public void start() {
|
||||
switch(method) {
|
||||
case MC_CONTROL_FIRST_VISIT:
|
||||
learning = new MonteCarloControlFirstVisitEGreedy<>(environment, discreteActionSpace, discountFactor, epsilon, delay);
|
||||
learning = new MonteCarloControlEGreedy<>(environment, discreteActionSpace, discountFactor, epsilon, delay);
|
||||
break;
|
||||
case SARSA_EPISODIC:
|
||||
case MC_CONTROL_EVERY_VISIT:
|
||||
learning = new MonteCarloControlEGreedy<>(environment, discreteActionSpace, discountFactor, epsilon, delay, true);
|
||||
break;
|
||||
|
||||
case SARSA_ON_POLICY_CONTROL:
|
||||
learning = new SARSA<>(environment, discreteActionSpace, discountFactor, epsilon, learningRate, delay);
|
||||
break;
|
||||
case Q_LEARNING_OFF_POLICY_CONTROL:
|
||||
@@ -67,17 +71,17 @@ public class RLController<A extends Enum> implements LearningListener {
|
||||
}
|
||||
|
||||
protected void initListeners() {
|
||||
learning.addListener(this);
|
||||
new Thread(() -> {
|
||||
while(true) {
|
||||
printNextEpisode = true;
|
||||
try {
|
||||
Thread.sleep(30 * 1000);
|
||||
} catch (InterruptedException e) {
|
||||
e.printStackTrace();
|
||||
learning.addListener(this);
|
||||
new Thread(() -> {
|
||||
while(learning.isCurrentlyLearning()) {
|
||||
printNextEpisode = true;
|
||||
try {
|
||||
Thread.sleep(30 * 1000);
|
||||
} catch (InterruptedException e) {
|
||||
e.printStackTrace();
|
||||
}
|
||||
}
|
||||
}
|
||||
}).start();
|
||||
}).start();
|
||||
}
|
||||
|
||||
private void initLearning() {
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
package core.policy;
|
||||
|
||||
/**
|
||||
* Chooses the action with the highest values with possibility: 1-Ɛ + Ɛ/|A|
|
||||
* With possibility of Ɛ, a random action is taken (highest values option included).
|
||||
* Chooses the action with the highest values with possibility: 1-Epsilon + Epsilon/|A|
|
||||
* With possibility of Epsilon, a random action is taken (highest values option included).
|
||||
*
|
||||
* @param <A> Enum class of available action in specific environment
|
||||
*/
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
package evironment.antGame;
|
||||
|
||||
import core.State;
|
||||
import lombok.AllArgsConstructor;
|
||||
|
||||
import java.util.Objects;
|
||||
|
||||
@AllArgsConstructor
|
||||
public class AntStateOriginal implements State {
|
||||
private final int currentFood;
|
||||
private final int row;
|
||||
private final int col;
|
||||
private final CellType type;
|
||||
private final int smell;
|
||||
private final int food;
|
||||
|
||||
@Override
|
||||
public boolean equals(Object o) {
|
||||
if (this == o) return true;
|
||||
if (o == null || getClass() != o.getClass()) return false;
|
||||
AntStateOriginal that = (AntStateOriginal) o;
|
||||
return currentFood == that.currentFood &&
|
||||
row == that.row &&
|
||||
col == that.col &&
|
||||
smell == that.smell &&
|
||||
type == that.type &&
|
||||
food == that.food;
|
||||
}
|
||||
|
||||
@Override
|
||||
public int hashCode() {
|
||||
return Objects.hash(currentFood, row, col, type, smell, food);
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "AntStateOriginal{" +
|
||||
"currentFood=" + currentFood +
|
||||
", row=" + row +
|
||||
", col=" + col +
|
||||
", type=" + type +
|
||||
", smell=" + smell +
|
||||
", food=" + food +
|
||||
'}';
|
||||
}
|
||||
}
|
||||
@@ -9,16 +9,19 @@ import evironment.antGame.gui.AntWorldComponent;
|
||||
import javax.swing.*;
|
||||
import java.awt.*;
|
||||
|
||||
/**
|
||||
* Episodic AntWorld
|
||||
*/
|
||||
public class AntWorld implements Environment<AntAction>, Visualizable {
|
||||
/**
|
||||
*
|
||||
*/
|
||||
private Grid grid;
|
||||
protected Grid grid;
|
||||
/**
|
||||
* Intern (backend) representation of the ant.
|
||||
* The AntWorld essentially acts like the game host of the original AntGame.
|
||||
*/
|
||||
private Ant myAnt;
|
||||
protected Ant myAnt;
|
||||
/**
|
||||
* The client agent. In the original AntGame the host would send jade messages
|
||||
* of the current observation to each client on every tick.
|
||||
@@ -32,13 +35,13 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
||||
* through an intern grid clone (brain), for example. A history as mentioned in
|
||||
* various lectures could be possible as well.
|
||||
*/
|
||||
private AntAgent antAgent;
|
||||
protected AntAgent antAgent;
|
||||
|
||||
private int tick;
|
||||
protected int tick;
|
||||
private int maxEpisodeTicks;
|
||||
|
||||
public AntWorld(int width, int height, double foodDensity){
|
||||
grid = new Grid(width, height, foodDensity);
|
||||
public AntWorld(int width, int height) {
|
||||
grid = new Grid(width, height);
|
||||
antAgent = new AntAgent(width, height);
|
||||
myAnt = new Ant();
|
||||
maxEpisodeTicks = 1000;
|
||||
@@ -46,73 +49,68 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
||||
}
|
||||
|
||||
public AntWorld(){
|
||||
this(Constants.DEFAULT_GRID_WIDTH, Constants.DEFAULT_GRID_HEIGHT, Constants.DEFAULT_FOOD_DENSITY);
|
||||
this(Constants.DEFAULT_GRID_WIDTH, Constants.DEFAULT_GRID_HEIGHT);
|
||||
}
|
||||
|
||||
@Override
|
||||
public StepResultEnvironment step(AntAction action){
|
||||
AntObservation observation;
|
||||
State newState;
|
||||
double reward = 0;
|
||||
String info = "";
|
||||
boolean done = false;
|
||||
|
||||
protected StepCalculation processStep(AntAction action) {
|
||||
StepCalculation sc = new StepCalculation();
|
||||
sc.reward = Reward.DEFAULT_REWARD;
|
||||
sc.info = "";
|
||||
sc.done = false;
|
||||
Cell currentCell = grid.getCell(myAnt.getPos());
|
||||
Point potentialNextPos = new Point(myAnt.getPos().x, myAnt.getPos().y);
|
||||
boolean stayOnCell = true;
|
||||
sc.potentialNextPos = new Point(myAnt.getPos().x, myAnt.getPos().y);
|
||||
sc.stayOnCell = true;
|
||||
// flag to enable a check if all food has been collected only fired if food was dropped
|
||||
// on the starting position
|
||||
boolean checkCompletion = false;
|
||||
sc.checkCompletion = false;
|
||||
|
||||
switch (action) {
|
||||
switch(action) {
|
||||
case MOVE_UP:
|
||||
potentialNextPos.y -= 1;
|
||||
stayOnCell = false;
|
||||
sc.potentialNextPos.y -= 1;
|
||||
sc.stayOnCell = false;
|
||||
break;
|
||||
case MOVE_RIGHT:
|
||||
potentialNextPos.x += 1;
|
||||
stayOnCell = false;
|
||||
sc.potentialNextPos.x += 1;
|
||||
sc.stayOnCell = false;
|
||||
break;
|
||||
case MOVE_DOWN:
|
||||
potentialNextPos.y += 1;
|
||||
stayOnCell = false;
|
||||
sc.potentialNextPos.y += 1;
|
||||
sc.stayOnCell = false;
|
||||
break;
|
||||
case MOVE_LEFT:
|
||||
potentialNextPos.x -= 1;
|
||||
stayOnCell = false;
|
||||
sc.potentialNextPos.x -= 1;
|
||||
sc.stayOnCell = false;
|
||||
break;
|
||||
case PICK_UP:
|
||||
if(myAnt.hasFood()){
|
||||
if(myAnt.hasFood()) {
|
||||
// Ant tries to pick up food but can only hold one piece
|
||||
reward = Reward.FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY;
|
||||
}else if(currentCell.getFood() == 0){
|
||||
sc.reward = Reward.FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY;
|
||||
} else if(currentCell.getFood() == 0) {
|
||||
// Ant tries to pick up food on cell that has no food on it
|
||||
reward = Reward.FOOD_PICK_UP_FAIL_NO_FOOD;
|
||||
}else if(currentCell.getFood() > 0){
|
||||
sc.reward = Reward.FOOD_PICK_UP_FAIL_NO_FOOD;
|
||||
} else if(currentCell.getFood() > 0) {
|
||||
// Ant successfully picks up food
|
||||
currentCell.setFood(currentCell.getFood() - 1);
|
||||
myAnt.setHasFood(true);
|
||||
reward = Reward.FOOD_PICK_UP_SUCCESS;
|
||||
sc.reward = Reward.FOOD_PICK_UP_SUCCESS;
|
||||
}
|
||||
break;
|
||||
case DROP_DOWN:
|
||||
if(!myAnt.hasFood()){
|
||||
if(!myAnt.hasFood()) {
|
||||
// Ant had no food to drop
|
||||
reward = Reward.FOOD_DROP_DOWN_FAIL_NO_FOOD;
|
||||
}else{
|
||||
// Drop food onto the ground
|
||||
currentCell.setFood(currentCell.getFood() + 1);
|
||||
sc.reward = Reward.FOOD_DROP_DOWN_FAIL_NO_FOOD;
|
||||
} else {
|
||||
myAnt.setHasFood(false);
|
||||
|
||||
// negative reward if the agent drops food on any other field
|
||||
// than the starting point
|
||||
if(currentCell.getType() != CellType.START){
|
||||
reward = Reward.FOOD_DROP_DOWN_FAIL_NOT_START;
|
||||
done = true;
|
||||
}else{
|
||||
reward = Reward.FOOD_DROP_DOWN_SUCCESS;
|
||||
if(currentCell.getType() != CellType.START) {
|
||||
sc.reward = Reward.FOOD_DROP_DOWN_FAIL_NOT_START;
|
||||
// Drop food onto the ground
|
||||
currentCell.setFood(currentCell.getFood() + 1);
|
||||
} else {
|
||||
sc.reward = Reward.FOOD_DROP_DOWN_SUCCESS;
|
||||
myAnt.setPoints(myAnt.getPoints() + 1);
|
||||
checkCompletion = true;
|
||||
sc.checkCompletion = true;
|
||||
}
|
||||
}
|
||||
break;
|
||||
@@ -121,60 +119,73 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
||||
}
|
||||
|
||||
// movement action was selected
|
||||
if(!stayOnCell){
|
||||
if(!isInGrid(potentialNextPos)){
|
||||
stayOnCell = true;
|
||||
reward = Reward.RAN_INTO_WALL;
|
||||
}else if(hitObstacle(potentialNextPos)){
|
||||
stayOnCell = true;
|
||||
reward = Reward.RAN_INTO_OBSTACLE;
|
||||
if(!sc.stayOnCell) {
|
||||
if(!isInGrid(sc.potentialNextPos)) {
|
||||
sc.stayOnCell = true;
|
||||
sc.reward = Reward.RAN_INTO_WALL;
|
||||
} else if(hitObstacle(sc.potentialNextPos)) {
|
||||
sc.stayOnCell = true;
|
||||
sc.reward = Reward.RAN_INTO_OBSTACLE;
|
||||
}
|
||||
}
|
||||
|
||||
// valid movement
|
||||
if(!stayOnCell){
|
||||
myAnt.getPos().setLocation(potentialNextPos);
|
||||
if(antAgent.getCell(myAnt.getPos()).getType() == CellType.UNKNOWN){
|
||||
// the ant will move to a cell that was previously unknown
|
||||
reward = Reward.UNKNOWN_FIELD_EXPLORED;
|
||||
}else{
|
||||
reward = 0;
|
||||
}
|
||||
}
|
||||
|
||||
// get observation after action was computed
|
||||
observation = new AntObservation(grid.getCell(myAnt.getPos()), myAnt.getPos(), myAnt.hasFood());
|
||||
|
||||
// let the ant agent process the observation to create a valid markov state
|
||||
newState = antAgent.feedObservation(observation);
|
||||
|
||||
if(checkCompletion){
|
||||
done = grid.isAllFoodCollected();
|
||||
}
|
||||
|
||||
|
||||
/*
|
||||
if(!done){
|
||||
reward = -1;
|
||||
}
|
||||
*/
|
||||
if(++tick == maxEpisodeTicks){
|
||||
done = true;
|
||||
}
|
||||
|
||||
|
||||
StepResultEnvironment result = new StepResultEnvironment(newState, reward, done, info);
|
||||
return result;
|
||||
return sc;
|
||||
}
|
||||
|
||||
private boolean isInGrid(Point pos){
|
||||
@Override
|
||||
public StepResultEnvironment step(AntAction action){
|
||||
StepCalculation sc = processStep(action);
|
||||
|
||||
// valid movement
|
||||
if(!sc.stayOnCell) {
|
||||
myAnt.getPos().setLocation(sc.potentialNextPos);
|
||||
if(antAgent.getCell(myAnt.getPos()).getType() == CellType.UNKNOWN){
|
||||
// the ant will move to a cell that was previously unknown
|
||||
// TODO: not optimal for going straight for food
|
||||
// sc.reward = Reward.UNKNOWN_FIELD_EXPLORED;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
if(sc.checkCompletion) {
|
||||
sc.done = grid.isAllFoodCollected();
|
||||
}
|
||||
|
||||
if(++tick == maxEpisodeTicks){
|
||||
sc.done = true;
|
||||
}
|
||||
|
||||
return new StepResultEnvironment(generateReturnState(), sc.reward, sc.done, sc.info);
|
||||
}
|
||||
|
||||
protected State generateReturnState(){
|
||||
// get observation after action was computed
|
||||
AntObservation observation = new AntObservation(grid.getCell(myAnt.getPos()), myAnt.getPos(), myAnt.hasFood());
|
||||
|
||||
// let the ant agent process the observation to create a valid markov state
|
||||
return antAgent.feedObservation(observation);
|
||||
}
|
||||
|
||||
protected boolean isInGrid(Point pos) {
|
||||
return pos.x >= 0 && pos.x < grid.getWidth() && pos.y >= 0 && pos.y < grid.getHeight();
|
||||
}
|
||||
|
||||
private boolean hitObstacle(Point pos){
|
||||
protected boolean hitObstacle(Point pos) {
|
||||
return grid.getCell(pos).getType() == CellType.OBSTACLE;
|
||||
}
|
||||
|
||||
protected class StepCalculation {
|
||||
double reward;
|
||||
String info;
|
||||
boolean done;
|
||||
Point potentialNextPos = new Point(myAnt.getPos().x, myAnt.getPos().y);
|
||||
boolean stayOnCell = true;
|
||||
// flag to enable a check if all food has been collected only fired if food was dropped
|
||||
// on the starting position
|
||||
boolean checkCompletion = false;
|
||||
}
|
||||
|
||||
public State reset() {
|
||||
grid.resetWorld();
|
||||
antAgent.initUnknownWorld();
|
||||
@@ -189,6 +200,7 @@ public class AntWorld implements Environment<AntAction>, Visualizable {
|
||||
public void setMaxEpisodeLength(int maxTicks){
|
||||
this.maxEpisodeTicks = maxTicks;
|
||||
}
|
||||
|
||||
public Point getSpawningPoint(){
|
||||
return grid.getStartPoint();
|
||||
}
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
package evironment.antGame;
|
||||
|
||||
import core.State;
|
||||
import core.StepResultEnvironment;
|
||||
|
||||
public class AntWorldContinuous extends AntWorld {
|
||||
public AntWorldContinuous(int width, int height) {
|
||||
super(width, height);
|
||||
}
|
||||
|
||||
public AntWorldContinuous() {
|
||||
super();
|
||||
}
|
||||
|
||||
@Override
|
||||
public StepResultEnvironment step(AntAction action) {
|
||||
Cell currentCell = grid.getCell(myAnt.getPos());
|
||||
|
||||
StepCalculation sc = processStep(action);
|
||||
|
||||
// flag is set to true if food gets dropped onto starts
|
||||
if(sc.checkCompletion) {
|
||||
grid.spawnNewFood();
|
||||
}
|
||||
// valid movement
|
||||
if(!sc.stayOnCell) {
|
||||
myAnt.getPos().setLocation(sc.potentialNextPos);
|
||||
}
|
||||
|
||||
return new StepResultEnvironment(generateReturnState(), sc.reward, false, sc.info);
|
||||
}
|
||||
|
||||
@Override
|
||||
protected State generateReturnState(){
|
||||
AntObservation observation = new AntObservation(grid.getCell(myAnt.getPos()), myAnt.getPos(), myAnt.hasFood());
|
||||
return new AntState(grid.getGrid(), observation.getPos(), observation.hasFood());
|
||||
}
|
||||
|
||||
}
|
||||
@@ -0,0 +1,36 @@
|
||||
package evironment.antGame;
|
||||
|
||||
import core.State;
|
||||
|
||||
public class AntWorldContinuousOriginalState extends AntWorldContinuous {
|
||||
public AntWorldContinuousOriginalState(int width, int height) {
|
||||
super(width, height);
|
||||
}
|
||||
|
||||
public AntWorldContinuousOriginalState() {
|
||||
super();
|
||||
}
|
||||
|
||||
@Override
|
||||
protected State generateReturnState(){
|
||||
return new AntStateOriginal(myAnt.hasFood()? 1:0, myAnt.getPos().x, myAnt.getPos().y, grid.getCell(myAnt.getPos()).getType(), calculateSmell(), grid.getCell(myAnt.getPos()).getFood());
|
||||
}
|
||||
|
||||
/**
|
||||
* @return total smell of neighbour food cells
|
||||
*/
|
||||
private int calculateSmell(){
|
||||
int smell = 0;
|
||||
int maxX = grid.getGrid().length -1;
|
||||
int maxY = grid.getGrid()[0].length -1;
|
||||
int antX = myAnt.getPos().x;
|
||||
int antY = myAnt.getPos().y;
|
||||
|
||||
smell += antY > 0 ? grid.getCell(antX, antY - 1).getFood() : 0;
|
||||
smell += antY < maxY ? grid.getCell(antX, antY + 1).getFood() : 0;
|
||||
smell += antX > 0 ? grid.getCell(antX - 1, antY).getFood() : 0;
|
||||
smell += antX < maxX ? grid.getCell(antX + 1, antY).getFood() : 0;
|
||||
|
||||
return smell;
|
||||
}
|
||||
}
|
||||
@@ -7,6 +7,7 @@ import java.awt.*;
|
||||
|
||||
public class Cell {
|
||||
@Getter
|
||||
@Setter
|
||||
private CellType type;
|
||||
@Getter
|
||||
@Setter
|
||||
@@ -38,4 +39,13 @@ public class Cell {
|
||||
}
|
||||
return super.equals(obj);
|
||||
}
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "Cell{" +
|
||||
"type=" + type +
|
||||
", food=" + food +
|
||||
", pos=" + pos +
|
||||
'}';
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package evironment.antGame;
|
||||
|
||||
public enum CellType {
|
||||
|
||||
public enum CellType{
|
||||
START,
|
||||
FREE,
|
||||
OBSTACLE,
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
package evironment.antGame;
|
||||
|
||||
public class Constants {
|
||||
public static final int DEFAULT_GRID_WIDTH = 10;
|
||||
public static final int DEFAULT_GRID_HEIGHT = 10;
|
||||
public static final double DEFAULT_FOOD_DENSITY = 0.1;
|
||||
public static final int DEFAULT_GRID_WIDTH = 5;
|
||||
public static final int DEFAULT_GRID_HEIGHT = 5;
|
||||
}
|
||||
|
||||
@@ -7,24 +7,18 @@ import java.awt.*;
|
||||
public class Grid {
|
||||
private int width;
|
||||
private int height;
|
||||
private double foodDensity;
|
||||
private Point start;
|
||||
private Cell[][] grid;
|
||||
private Cell[][] initialGrid;
|
||||
|
||||
public Grid(int width, int height, double foodDensity){
|
||||
public Grid(int width, int height) {
|
||||
this.width = width;
|
||||
this.height = height;
|
||||
this.foodDensity = foodDensity;
|
||||
grid = new Cell[width][height];
|
||||
initialGrid = new Cell[width][height];
|
||||
initRandomWorld();
|
||||
}
|
||||
|
||||
public Grid(int width, int height){
|
||||
this(width, height, 0);
|
||||
}
|
||||
|
||||
public void resetWorld(){
|
||||
grid = Util.deepCopyCellGrid(initialGrid);
|
||||
}
|
||||
@@ -32,15 +26,52 @@ public class Grid {
|
||||
public void initRandomWorld(){
|
||||
for(int x = 0; x < width; ++x){
|
||||
for(int y = 0; y < height; ++y){
|
||||
if( RNG.getRandom().nextDouble() < foodDensity){
|
||||
initialGrid[x][y] = new Cell(new Point(x,y), CellType.FREE, 1);
|
||||
}else{
|
||||
initialGrid[x][y] = new Cell(new Point(x,y), CellType.FREE);
|
||||
}
|
||||
initialGrid[x][y] = new Cell(new Point(x, y), CellType.FREE);
|
||||
}
|
||||
}
|
||||
start = new Point(RNG.getRandom().nextInt(width), RNG.getRandom().nextInt(height));
|
||||
start = new Point(RNG.getRandomEnv().nextInt(width), RNG.getRandomEnv().nextInt(height));
|
||||
initialGrid[start.x][start.y] = new Cell(new Point(start.x, start.y), CellType.START);
|
||||
spawnNewFood(initialGrid);
|
||||
spawnObstacles();
|
||||
}
|
||||
|
||||
//TODO
|
||||
private void spawnObstacles() {
|
||||
initialGrid[3][1].setType(CellType.OBSTACLE);
|
||||
initialGrid[4][1].setType(CellType.OBSTACLE);
|
||||
initialGrid[5][1].setType(CellType.OBSTACLE);
|
||||
initialGrid[6][1].setType(CellType.OBSTACLE);
|
||||
initialGrid[7][1].setType(CellType.OBSTACLE);
|
||||
initialGrid[3][2].setType(CellType.OBSTACLE);
|
||||
initialGrid[3][3].setType(CellType.OBSTACLE);
|
||||
initialGrid[3][4].setType(CellType.OBSTACLE);
|
||||
initialGrid[4][4].setType(CellType.OBSTACLE);
|
||||
initialGrid[5][4].setType(CellType.OBSTACLE);
|
||||
initialGrid[6][4].setType(CellType.OBSTACLE);
|
||||
}
|
||||
|
||||
/**
|
||||
* Spawns one additional food on a random field EXCEPT for the starting position
|
||||
*/
|
||||
public void spawnNewFood(Cell[][] grid) {
|
||||
boolean foodSpawned = false;
|
||||
Point potFood = new Point(0, 0);
|
||||
CellType potFieldType;
|
||||
while(!foodSpawned) {
|
||||
potFood.x = RNG.getRandomEnv().nextInt(width);
|
||||
potFood.y = RNG.getRandomEnv().nextInt(height);
|
||||
potFieldType = grid[potFood.x][potFood.y].getType();
|
||||
if(potFieldType != CellType.START && grid[potFood.x][potFood.y].getFood() == 0 && potFieldType != CellType.OBSTACLE) {
|
||||
grid[potFood.x][potFood.y].setFood(1);
|
||||
foodSpawned = true;
|
||||
// System.out.println("spawned new food at " + potFood);
|
||||
// System.out.println(initialGrid[potFood.x][potFood.y]);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
public void spawnNewFood() {
|
||||
spawnNewFood(grid);
|
||||
}
|
||||
|
||||
public Point getStartPoint(){
|
||||
|
||||
@@ -1,16 +1,17 @@
|
||||
package evironment.antGame;
|
||||
|
||||
public class Reward {
|
||||
public static final double FOOD_PICK_UP_SUCCESS = 1;
|
||||
public static final double FOOD_PICK_UP_FAIL_NO_FOOD = -1;
|
||||
public static final double FOOD_PICK_UP_FAIL_HAS_FOOD_ALREADY = -1;
|
||||
public static final double DEFAULT_REWARD = -1;
|
||||
public static final double FOOD_PICK_UP_SUCCESS = 0;
|
||||
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_NOT_START = -1;
|
||||
public static final double FOOD_DROP_DOWN_FAIL_NO_FOOD = -2;
|
||||
public static final double FOOD_DROP_DOWN_FAIL_NOT_START = -2;
|
||||
public static final double FOOD_DROP_DOWN_SUCCESS = 1;
|
||||
|
||||
public static final double UNKNOWN_FIELD_EXPLORED = 1;
|
||||
public static final double UNKNOWN_FIELD_EXPLORED = 0;
|
||||
|
||||
public static final double RAN_INTO_WALL = -1;
|
||||
public static final double RAN_INTO_OBSTACLE = -1;
|
||||
public static final double RAN_INTO_WALL = -2;
|
||||
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
|
||||
and EXCLUSIVE! bound
|
||||
*/
|
||||
return cards.get(RNG.getRandom().nextInt(cards.size()));
|
||||
return cards.get(RNG.getRandomEnv().nextInt(cards.size()));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
package evironment.jumpingDino;
|
||||
|
||||
import core.State;
|
||||
import core.gui.Visualizable;
|
||||
import lombok.AllArgsConstructor;
|
||||
import lombok.Getter;
|
||||
|
||||
import javax.swing.*;
|
||||
import java.awt.*;
|
||||
import java.io.Serializable;
|
||||
import java.util.Objects;
|
||||
|
||||
@AllArgsConstructor
|
||||
@Getter
|
||||
public class DinoStateSimple implements State, Serializable, Visualizable {
|
||||
protected final double scale = 0.5;
|
||||
private int xDistanceToObstacle;
|
||||
|
||||
@Override
|
||||
public String toString() {
|
||||
return "DinoState{" +
|
||||
"xDistanceToObstacle=" + xDistanceToObstacle +
|
||||
'}';
|
||||
}
|
||||
|
||||
@Override
|
||||
public boolean equals(Object o) {
|
||||
if(this == o) return true;
|
||||
if(o == null || getClass() != o.getClass()) return false;
|
||||
DinoStateSimple dinoState = (DinoStateSimple) o;
|
||||
return xDistanceToObstacle == dinoState.xDistanceToObstacle;
|
||||
}
|
||||
|
||||
@Override
|
||||
public int hashCode() {
|
||||
return Objects.hash(xDistanceToObstacle);
|
||||
}
|
||||
|
||||
@Override
|
||||
public JComponent visualize() {
|
||||
return new JComponent() {
|
||||
{
|
||||
setPreferredSize(new Dimension(Config.FRAME_WIDTH, (int) (scale * Config.FRAME_HEIGHT)));
|
||||
setVisible(true);
|
||||
}
|
||||
|
||||
@Override
|
||||
protected void paintComponent(Graphics g) {
|
||||
super.paintComponents(g);
|
||||
drawObjects(g);
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
public void drawObjects(Graphics g) {
|
||||
g.setColor(Color.BLACK);
|
||||
g.fillRect(0, (int) (scale * (Config.FRAME_HEIGHT - Config.GROUND_Y)), Config.FRAME_WIDTH, 2);
|
||||
|
||||
g.fillRect((int) (scale * Config.DINO_STARTING_X), (int) (scale * (Config.FRAME_HEIGHT - Config.GROUND_Y - Config.DINO_SIZE)), (int) (scale * Config.DINO_SIZE), (int) (scale * Config.DINO_SIZE));
|
||||
g.drawString("Distance: " + xDistanceToObstacle, (int) (scale * Config.DINO_STARTING_X), (int) (scale * (Config.FRAME_HEIGHT - Config.GROUND_Y - Config.OBSTACLE_SIZE - 40)));
|
||||
|
||||
g.fillRect((int) (scale * (Config.DINO_STARTING_X + getXDistanceToObstacle())), (int) (scale * (Config.FRAME_HEIGHT - Config.GROUND_Y - Config.OBSTACLE_SIZE)), (int) (scale * Config.OBSTACLE_SIZE), (int) (scale * Config.OBSTACLE_SIZE));
|
||||
|
||||
}
|
||||
}
|
||||
@@ -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;
|
||||
@@ -46,19 +49,6 @@ public class DinoWorld implements Environment<DinoAction>, Visualizable {
|
||||
if(action == DinoAction.JUMP){
|
||||
dino.jump();
|
||||
}
|
||||
|
||||
// for(int i= 0; i < 5; ++i){
|
||||
// dino.tick();
|
||||
// currentObstacle.tick();
|
||||
// if(currentObstacle.getX() < -Config.OBSTACLE_SIZE){
|
||||
// spawnNewObstacle();
|
||||
// }
|
||||
// comp.repaint();
|
||||
// if(ranIntoObstacle()){
|
||||
// done = true;
|
||||
// break;
|
||||
// }
|
||||
// }
|
||||
dino.tick();
|
||||
currentObstacle.tick();
|
||||
if(currentObstacle.getX() < -Config.OBSTACLE_SIZE) {
|
||||
@@ -72,6 +62,9 @@ public class DinoWorld implements Environment<DinoAction>, Visualizable {
|
||||
return new StepResultEnvironment(generateReturnState(), reward, done, "");
|
||||
}
|
||||
|
||||
protected State generateReturnState(){
|
||||
return new DinoStateSimple(getDistanceToObstacle());
|
||||
}
|
||||
protected State generateReturnState(){
|
||||
return new DinoState(getDistanceToObstacle(), dino.isInJump());
|
||||
}
|
||||
|
||||
@@ -1,9 +1,22 @@
|
||||
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
|
||||
*
|
||||
* 2044
|
||||
* 4 speeds, 4 distance
|
||||
*/
|
||||
public class DinoWorldAdvanced extends DinoWorld{
|
||||
public DinoWorldAdvanced(){
|
||||
super();
|
||||
@@ -11,16 +24,35 @@ public class DinoWorldAdvanced extends DinoWorld{
|
||||
|
||||
@Override
|
||||
protected State generateReturnState() {
|
||||
return new DinoStateWithSpeed(getDistanceToObstacle(), dino.isInJump(), getCurrentObstacle().getDx());
|
||||
return new DinoStateWithSpeed(getDistanceToObstacle(), dino.isInJump(), currentObstacle.getDx());
|
||||
}
|
||||
|
||||
@Override
|
||||
protected void spawnNewObstacle() {
|
||||
int dx;
|
||||
int xSpawn;
|
||||
dx = -(int)((Math.random() + 0.5) * 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);
|
||||
double ran = RNG.getRandomEnv().nextDouble();
|
||||
if(ran < 0.25){
|
||||
dx = -(int) (0.35 * Config.OBSTACLE_SPEED);
|
||||
}else if(ran < 0.5){
|
||||
dx = -(int) (0.7 * Config.OBSTACLE_SPEED);
|
||||
}else if(ran < 0.75){
|
||||
dx = -(int)(1.6 * Config.OBSTACLE_SPEED);
|
||||
} else{
|
||||
dx = -(int) (3.5 * Config.OBSTACLE_SPEED);
|
||||
}
|
||||
double ran2 = RNG.getRandomEnv().nextDouble();
|
||||
if(ran2 < 0.25) {
|
||||
// randomly spawning more right outside of the screen
|
||||
xSpawn = Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE;
|
||||
|
||||
} else if(ran2 < 0.5) {
|
||||
xSpawn = (int) (1.08 * Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE);
|
||||
} else if(ran2 < 0.75) {
|
||||
xSpawn = (int) (1.11 * Config.FRAME_WIDTH + Config.FRAME_WIDTH + Config.OBSTACLE_SIZE);
|
||||
} else {
|
||||
xSpawn = (int) (1.23 * 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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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();
|
||||
}
|
||||
}
|
||||
@@ -3,24 +3,24 @@ package example;
|
||||
import core.RNG;
|
||||
import core.algo.Method;
|
||||
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);
|
||||
RNG.setSeed(29);
|
||||
|
||||
RLController<DinoAction> rl = new RLControllerGUI<>(
|
||||
RLController<DinoAction> rl = new RLController<>(
|
||||
new DinoWorldAdvanced(),
|
||||
Method.MC_CONTROL_FIRST_VISIT,
|
||||
DinoAction.values());
|
||||
|
||||
rl.setDelay(100);
|
||||
rl.setDiscountFactor(1f);
|
||||
rl.setEpsilon(0.15f);
|
||||
rl.setLearningRate(1f);
|
||||
rl.setNrOfEpisodes(50001);
|
||||
rl.setDelay(0);
|
||||
rl.setDiscountFactor(9f);
|
||||
rl.setEpsilon(0.05f);
|
||||
rl.setLearningRate(0.8f);
|
||||
rl.setNrOfEpisodes(100000);
|
||||
rl.start();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,15 +12,15 @@ public class RunningAnt {
|
||||
RNG.setSeed(56);
|
||||
|
||||
RLController<AntAction> rl = new RLControllerGUI<>(
|
||||
new AntWorld(3, 3, 0.1),
|
||||
Method.MC_CONTROL_FIRST_VISIT,
|
||||
new AntWorld(8, 8),
|
||||
Method.Q_LEARNING_OFF_POLICY_CONTROL,
|
||||
AntAction.values());
|
||||
|
||||
rl.setDelay(200);
|
||||
rl.setNrOfEpisodes(10000);
|
||||
rl.setDiscountFactor(1f);
|
||||
rl.setDiscountFactor(0.9f);
|
||||
rl.setLearningRate(0.9f);
|
||||
rl.setEpsilon(0.15f);
|
||||
|
||||
rl.start();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,52 +0,0 @@
|
||||
package example;
|
||||
|
||||
public class Test {
|
||||
interface Drawable{
|
||||
void draw();
|
||||
}
|
||||
interface State{
|
||||
int getInt();
|
||||
}
|
||||
|
||||
static class A implements Drawable, State{
|
||||
private int k;
|
||||
public A(int a){
|
||||
k = a;
|
||||
}
|
||||
@Override
|
||||
public void draw() {
|
||||
System.out.println("draw " + k);
|
||||
}
|
||||
|
||||
@Override
|
||||
public int getInt() {
|
||||
System.out.println("getInt" + k);
|
||||
return k;
|
||||
}
|
||||
}
|
||||
|
||||
static class B implements State{
|
||||
@Override
|
||||
public int getInt() {
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
public static void main(String[] args) {
|
||||
State state = new A(24);
|
||||
State state2 = new B();
|
||||
state.getInt();
|
||||
|
||||
System.out.println(state2 instanceof Drawable);
|
||||
drawState(state2);
|
||||
}
|
||||
|
||||
static void drawState(State s){
|
||||
if(s instanceof Drawable){
|
||||
Drawable d = (Drawable) s;
|
||||
d.draw();
|
||||
}else{
|
||||
System.out.println("invalid");
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user