Skip to content

Commit bdab26c

Browse files
Merge pull request dnv-opensource#4 from dnv-opensource/eis
Fixed the tests so that pytest can be run on the package. Added the p…
2 parents 1b8495b + dd5de9a commit bdab26c

16 files changed

Lines changed: 90 additions & 1761 deletions

scripts/play_q.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,9 @@
1515
def main():
1616
parser = argparse.ArgumentParser(description="Run a trained Q-learning agent on the crane anti-pendulum task.")
1717
parser.add_argument("--model-path", type=str, required=True, help="Path to a trained Q-table JSON")
18-
parser.add_argument("--render-mode", type=str, default="plot", help="Render mode (plot, play-back, reward-tracking)")
18+
parser.add_argument(
19+
"--render-mode", type=str, default="plot", help="Render mode (plot, play-back, reward-tracking)"
20+
)
1921
parser.add_argument("--episodes", type=int, default=1, help="Number of episodes to run")
2022
parser.add_argument("--v0", type=float, default=-1.0, help="Initial crane speed (negative = stop mode)")
2123
args = parser.parse_args()

src/crane_controller/algorithm.py

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -21,9 +21,6 @@ def __init__(
2121
):
2222
self.env = env
2323
assert type(self.env).__name__ in AlgorithmAgent.envs, f"Environment {type(self.env).__name__} not listed."
24-
# print("ACTION_SPACE.N", env.action_space.n, defaultdict( lambda: np.zeros(env.action_space.n))['xx'])
25-
# Q-table: maps (state, action) to expected reward
26-
# defaultdict automatically creates entries with zeros for new states
2724

2825
# Track learning progress
2926
self.training_error: list[float] = []
@@ -75,8 +72,9 @@ def do_strategies(self, max_steps: int = 5000):
7572
if steps > max_steps:
7673
truncated = True
7774
res.append(reward)
78-
for i, self.strategy in enumerate(product(range(3), range(3), range(3), range(3))):
79-
print(f"{i}. strategy {self.strategy}: {res[i]}")
75+
if not self.env.render_mode == 'none':
76+
for i, self.strategy in enumerate(product(range(3), range(3), range(3), range(3))):
77+
print(f"{i}. strategy {self.strategy}: {res[i]}")
8078

8179
def do_episodes(self, n_episodes: int = 1000, show: int = 0, max_steps: int = 1000):
8280
"""Run episodes."""

src/crane_controller/envs/controlled_crane_pendulum.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ def __init__(
7373
elif render_mode == "plot":
7474
self.traces: dict[str, list[float]] = {"c_x": [], "c_v": [], "l_x": [], "l_v": []}
7575

76-
self.obeservation_space : spaces.Box | spaces.Discrete
76+
self.obeservation_space: spaces.Box | spaces.Discrete
7777
# Observations is a 4-dim np-array with
7878
# (crane-x, crane-v_x, load-polar-angle_x, load-v_x)
7979
self.min_speed = 0.1 # np.sqrt(2*reward_limit) # starting with less does not make sense (goal already reached)
@@ -98,7 +98,8 @@ def __init__(
9898
self.dt = dt
9999

100100
# We have 1 acceleration action which can each be min, zero or max, corresponding to acceleration of crane
101-
self.action_space = spaces.Discrete(3, start=0, seed=42, dtype=np.int64)
101+
super().reset(seed=seed) # make sure that the environment seed is set
102+
self.action_space = spaces.Discrete(3, start=0, dtype=np.int64)
102103
self.action_to_acc = {0: -self.acc, 1: 0.0, 2: self.acc}
103104

104105
def _init_discrete(self, spec: dict[str, tuple[float, ...]]):
@@ -211,7 +212,7 @@ def level(idx: int, val: float, categories: tuple[float, ...]):
211212
# if the crane moves towards the origo we do not add 'energy'
212213
self.reward = reward
213214

214-
obs : tuple[int,...] | np.ndarray
215+
obs: tuple[int, ...] | np.ndarray
215216
if len(self.discrete):
216217
obs = (
217218
level(1, energy, self.discrete["energies"]), # energy level

0 commit comments

Comments
 (0)