adopt MVC pattern and add real time graph interface

This commit is contained in:
2019-12-18 16:48:24 +01:00
parent 7f18a66e98
commit e0160ca1df
18 changed files with 450 additions and 29 deletions
+1 -1
View File
@@ -3,7 +3,7 @@ package core;
import java.util.List;
public interface DiscreteActionSpace<A extends Enum> {
int getNumberOfAction();
int getNumberOfActions();
void addAction(A a);
void addActions(A... as);
List<A> getAllActions();
+7
View File
@@ -0,0 +1,7 @@
package core;
public class LearningConfig {
public static final int DEFAULT_DELAY = 1;
public static final float DEFAULT_EPSILON = 0.1f;
public static final float DEFAULT_DISCOUNT_FACTOR = 1.0f;
}
@@ -32,7 +32,7 @@ public class ListDiscreteActionSpace<A extends Enum> implements DiscreteActionSp
}
@Override
public int getNumberOfAction(){
public int getNumberOfActions(){
return actions.size();
}
}
+41 -5
View File
@@ -2,26 +2,62 @@ package core.algo;
import core.DiscreteActionSpace;
import core.Environment;
import core.LearningConfig;
import core.StateActionTable;
import core.listener.LearningListener;
import core.policy.Policy;
import lombok.Getter;
import lombok.Setter;
import javax.swing.*;
import java.util.HashSet;
import java.util.Set;
@Getter
public abstract class Learning<A extends Enum> {
protected Policy<A> policy;
protected DiscreteActionSpace<A> actionSpace;
protected StateActionTable<A> stateActionTable;
protected Environment<A> environment;
protected float discountFactor;
@Setter
protected float epsilon;
protected Set<LearningListener> learningListeners;
@Setter
protected int delay;
public Learning(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon){
public Learning(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon, int delay){
this.environment = environment;
this.actionSpace = actionSpace;
this.discountFactor = discountFactor;
this.epsilon = epsilon;
}
public Learning(Environment<A> environment, DiscreteActionSpace<A> actionSpace){
this(environment, actionSpace, 1.0f, 0.1f);
this.delay = delay;
learningListeners = new HashSet<>();
}
public abstract void learn(int nrOfEpisodes, int delay);
public Learning(Environment<A> environment, DiscreteActionSpace<A> actionSpace, float discountFactor, float epsilon){
this(environment, actionSpace, discountFactor, epsilon, LearningConfig.DEFAULT_DELAY);
}
public Learning(Environment<A> environment, DiscreteActionSpace<A> actionSpace){
this(environment, actionSpace, LearningConfig.DEFAULT_DISCOUNT_FACTOR, LearningConfig.DEFAULT_EPSILON, LearningConfig.DEFAULT_DELAY);
}
public abstract void learn(int nrOfEpisodes);
public void addListener(LearningListener learningListener){
learningListeners.add(learningListener);
}
protected void dispatchEpisodeEnd(double sum){
for(LearningListener l: learningListeners) {
l.onEpisodeEnd(sum);
}
}
protected void dispatchEpisodeStart(){
for(LearningListener l: learningListeners){
l.onEpisodeStart();
}
}
}
@@ -34,22 +34,21 @@ public class MonteCarloOnPolicyEGreedy<A extends Enum> extends Learning<A> {
}
@Override
public void learn(int nrOfEpisodes, int delay) {
public void learn(int nrOfEpisodes) {
Map<Pair<State, A>, Double> returnSum = new HashMap<>();
Map<Pair<State, A>, Integer> returnCount = new HashMap<>();
State startingState = environment.reset();
for(int i = 0; i < nrOfEpisodes; ++i) {
List<StepResult<A>> episode = new ArrayList<>();
State state = environment.reset();
double rewardSum = 0;
double sumOfRewards = 0;
for(int j=0; j < 10; ++j){
Map<A, Double> actionValues = stateActionTable.getActionValues(state);
A chosenAction = policy.chooseAction(actionValues);
StepResultEnvironment envResult = environment.step(chosenAction);
State nextState = envResult.getState();
rewardSum += envResult.getReward();
sumOfRewards += envResult.getReward();
episode.add(new StepResult<>(state, chosenAction, envResult.getReward()));
if(envResult.isDone()) break;
@@ -57,13 +56,14 @@ public class MonteCarloOnPolicyEGreedy<A extends Enum> extends Learning<A> {
state = nextState;
try {
Thread.sleep(1);
Thread.sleep(delay);
} catch (InterruptedException e) {
e.printStackTrace();
}
}
System.out.printf("Episode %d \t Reward: %f \n", i, rewardSum);
dispatchEpisodeEnd(sumOfRewards);
System.out.printf("Episode %d \t Reward: %f \n", i, sumOfRewards);
Set<Pair<State, A>> stateActionPairs = new HashSet<>();
for(StepResult<A> sr: episode){
+5
View File
@@ -0,0 +1,5 @@
package core.algo;
public enum Method {
MC_ONPOLICY_EGREEDY, TD_ONPOLICY
}
@@ -0,0 +1,81 @@
package core.controller;
import core.DiscreteActionSpace;
import core.Environment;
import core.ListDiscreteActionSpace;
import core.algo.Learning;
import core.algo.Method;
import core.algo.mc.MonteCarloOnPolicyEGreedy;
import core.gui.View;
import javax.swing.*;
import java.util.Optional;
public class RLController<A extends Enum> implements ViewListener{
protected Environment<A> environment;
protected Learning<A> learning;
protected DiscreteActionSpace<A> discreteActionSpace;
protected View<A> view;
private int delay;
private int nrOfEpisodes;
private Method method;
public RLController(){
}
public void start(){
if(environment == null || discreteActionSpace == null || method == null){
throw new RuntimeException("Set environment, discreteActionSpace and method before calling .start()");
}
switch (method){
case MC_ONPOLICY_EGREEDY:
learning = new MonteCarloOnPolicyEGreedy<>(environment, discreteActionSpace);
break;
case TD_ONPOLICY:
break;
default:
throw new RuntimeException("Undefined method");
}
SwingUtilities.invokeLater(() ->{
view = new View<>(learning, this);
learning.addListener(view);
});
learning.learn(nrOfEpisodes);
}
@Override
public void onEpsilonChange(float epsilon) {
learning.setEpsilon(epsilon);
SwingUtilities.invokeLater(() -> view.updateLearningInfoPanel());
}
@Override
public void onDelayChange(int delay) {
}
public RLController<A> setMethod(Method method){
this.method = method;
return this;
}
public RLController<A> setEnvironment(Environment<A> environment){
this.environment = environment;
return this;
}
@SafeVarargs
public final RLController<A> setAllowedActions(A... actions){
this.discreteActionSpace = new ListDiscreteActionSpace<>(actions);
return this;
}
public RLController<A> setDelay(int delay){
this.delay = delay;
return this;
}
public RLController<A> setEpisodes(int nrOfEpisodes){
this.nrOfEpisodes = nrOfEpisodes;
return this;
}
}
@@ -0,0 +1,6 @@
package core.controller;
public interface ViewListener {
void onEpsilonChange(float epsilon);
void onDelayChange(int delay);
}
@@ -0,0 +1,41 @@
package core.gui;
import core.algo.Learning;
import core.controller.ViewListener;
import javax.swing.*;
public class LearningInfoPanel extends JPanel {
private Learning learning;
private JLabel policyLabel;
private JLabel discountLabel;
private JLabel epsilonLabel;
private JSlider epsilonSlider;
private JSlider delaySlider;
public LearningInfoPanel(Learning learning, ViewListener viewListener){
this.learning = learning;
setLayout(new BoxLayout(this, BoxLayout.Y_AXIS));
policyLabel = new JLabel();
discountLabel = new JLabel();
epsilonLabel = new JLabel();
epsilonSlider = new JSlider(0, 100, (int)(learning.getEpsilon() * 100));
epsilonSlider.addChangeListener(e -> viewListener.onEpsilonChange(epsilonSlider.getValue() / 100f));
add(policyLabel);
add(discountLabel);
add(epsilonLabel);
add(epsilonSlider);
refreshLabels();
setVisible(true);
}
public void refreshLabels(){
policyLabel.setText("Policy: " + learning.getPolicy().getClass());
discountLabel.setText("Discount factor: " + learning.getDiscountFactor());
epsilonLabel.setText("Exploration (Epsilon): " + learning.getEpsilon());
}
protected JSlider getEpsilonSlider(){
return epsilonSlider;
}
}
+102
View File
@@ -0,0 +1,102 @@
package core.gui;
import core.algo.Learning;
import core.controller.ViewListener;
import core.listener.LearningListener;
import lombok.Getter;
import org.knowm.xchart.QuickChart;
import org.knowm.xchart.XChartPanel;
import org.knowm.xchart.XYChart;
import javax.swing.*;
import java.awt.*;
import java.util.ArrayList;
import java.util.List;
public class View<A extends Enum> implements LearningListener {
private Learning<A> learning;
@Getter
private XYChart chart;
@Getter
private LearningInfoPanel learningInfoPanel;
@Getter
private JFrame mainFrame;
private XChartPanel<XYChart> rewardChartPanel;
private ViewListener viewListener;
private List<Double> rewardHistory;
public View(Learning<A> learning, ViewListener viewListener){
this.learning = learning;
this.viewListener = viewListener;
rewardHistory = new ArrayList<>();
this.initMainFrame();
}
private void initMainFrame(){
mainFrame = new JFrame();
mainFrame.setPreferredSize(new Dimension(1280, 720));
mainFrame.setLayout(new BorderLayout());
initLearningInfoPanel();
initRewardChart();
mainFrame.add(BorderLayout.WEST, learningInfoPanel);
mainFrame.add(BorderLayout.CENTER, rewardChartPanel);
mainFrame.setDefaultCloseOperation(WindowConstants.EXIT_ON_CLOSE);
mainFrame.pack();
mainFrame.setVisible(true);
}
private void initLearningInfoPanel(){
learningInfoPanel = new LearningInfoPanel(learning, viewListener);
}
private void initRewardChart(){
chart =
QuickChart.getChart(
"Rewards per Episode",
"Episode",
"Reward",
"randomWalk",
new double[] {0},
new double[] {0});
chart.getStyler().setLegendVisible(true);
chart.getStyler().setXAxisTicksVisible(true);
rewardChartPanel = new XChartPanel<>(chart);
rewardChartPanel.setPreferredSize(new Dimension(300,300));
}
public void showState(Visualizable state){
new JFrame(){
{
JComponent stateComponent = state.visualize();
setPreferredSize(new Dimension(stateComponent.getWidth(), stateComponent.getHeight()));
add(stateComponent);
setVisible(true);
}
};
}
public void updateRewardGraph(double recentReward){
rewardHistory.add(recentReward);
chart.updateXYSeries("randomWalk", null, rewardHistory, null);
rewardChartPanel.revalidate();
rewardChartPanel.repaint();
}
public void updateLearningInfoPanel(){
this.learningInfoPanel.refreshLabels();
}
@Override
public void onEpisodeEnd(double sumOfRewards) {
SwingUtilities.invokeLater(()->updateRewardGraph(sumOfRewards));
}
@Override
public void onEpisodeStart() {
}
}
+7
View File
@@ -0,0 +1,7 @@
package core.gui;
import javax.swing.*;
public interface Visualizable {
JComponent visualize();
}
@@ -0,0 +1,6 @@
package core.listener;
public interface LearningListener{
void onEpisodeEnd(double sumOfRewards);
void onEpisodeStart();
}