hanabhi/gridworld-env
0
1from typing import Optional2from core.env_server import Environment3from core.env_server.interfaces import Transform4from models import GridAction, GridObservation, GridState5 6 7class GridWorldEnv(Environment[GridAction, GridObservation, GridState]):8 def __init__(9 self,10 transform: Optional[Transform[GridObservation]] = None,11 rubric=None,12 ):13 super().__init__(transform=transform, rubric=rubric)14 15 self.grid_size = 516 self.agent_pos = (0, 0)17 self.goal_pos = (4, 4)18 self._steps = 019 20 def reset(self) -> GridObservation:21 self.agent_pos = (0, 0)22 self._steps = 023 return self._make_obs(reward=0.0, done=False)24 25 def step(self, action: GridAction) -> GridObservation:26 self._steps += 127 r, c = self.agent_pos28 29 move = action.direction.lower()30 if move == "right" and c < self.grid_size - 1: c += 131 elif move == "left" and c > 0: c -= 132 elif move == "down" and r < self.grid_size - 1: r += 133 elif move == "up" and r > 0: r -= 134 35 self.agent_pos = (r, c)36 done = self.agent_pos == self.goal_pos37 reward = 10.0 if done else -1.038 return self._make_obs(reward=reward, done=done)39 40 def _make_obs(self, reward: float, done: bool) -> GridObservation:41 return GridObservation(42 agent_pos=self.agent_pos,43 goal_pos=self.goal_pos,44 reward=reward,45 done=done,46 )47 48 def get_state(self) -> GridState:49 return GridState(episode_id="env_01", steps_taken=self._steps)50 51 @property52 def state(self) -> GridState:53 return self.get_state()