|
3 | 3 | import logging |
4 | 4 | from abc import ABC, abstractmethod |
5 | 5 | from copy import deepcopy |
6 | | -from typing import TYPE_CHECKING, Any, Callable, Optional, Union |
| 6 | +from typing import TYPE_CHECKING, Any, Union |
7 | 7 |
|
8 | 8 | import numpy as np |
9 | | -from Basilisk.utilities import orbitalMotion |
10 | 9 | from gymnasium import spaces |
11 | 10 |
|
12 | | -from bsk_rl.utils.functional import vectorize_nested_dict |
| 11 | +from bsk_rl.utils.functional import Resetable, vectorize_nested_dict |
13 | 12 | from bsk_rl.utils.orbital import rv2HN |
14 | 13 |
|
15 | 14 | if TYPE_CHECKING: # pragma: no cover |
@@ -49,7 +48,7 @@ def nested_obs_to_space(obs_dict, dtype): |
49 | 48 | raise TypeError(f"Cannot convert {obs_dict} to gym space.") |
50 | 49 |
|
51 | 50 |
|
52 | | -class ObservationBuilder: |
| 51 | +class ObservationBuilder(Resetable): |
53 | 52 | def __init__( |
54 | 53 | self, |
55 | 54 | satellite: "Satellite", |
@@ -80,6 +79,21 @@ def __init__( |
80 | 79 | name_counts[obs.name] = 1 |
81 | 80 | obs.link_satellite(self.satellite) |
82 | 81 |
|
| 82 | + def reset_overwrite_previous(self) -> None: |
| 83 | + """Perform any once-per-episode setup.""" |
| 84 | + for obs in self.observation_spec: |
| 85 | + obs.reset_overwrite_previous() |
| 86 | + |
| 87 | + def reset_pre_sim_init(self) -> None: |
| 88 | + """Perform any once-per-episode setup.""" |
| 89 | + for obs in self.observation_spec: |
| 90 | + obs.reset_pre_sim_init() |
| 91 | + |
| 92 | + def reset_during_sim_init(self) -> None: |
| 93 | + """Perform any once-per-episode setup.""" |
| 94 | + for obs in self.observation_spec: |
| 95 | + obs.reset_during_sim_init() |
| 96 | + |
83 | 97 | def reset_post_sim_init(self) -> None: |
84 | 98 | """Perform any once-per-episode setup.""" |
85 | 99 | self.simulator = self.satellite.simulator # already a proxy |
@@ -143,7 +157,7 @@ def observation_description(self) -> Any: |
143 | 157 | return self.obs_array_keys() |
144 | 158 |
|
145 | 159 |
|
146 | | -class Observation(ABC): |
| 160 | +class Observation(ABC, Resetable): |
147 | 161 | """Base observations class.""" |
148 | 162 |
|
149 | 163 | def __init__(self, name: str = "obs") -> None: |
@@ -176,10 +190,6 @@ def link_simulator(self, simulator: "Simulator") -> None: |
176 | 190 | """ |
177 | 191 | self.simulator = simulator # already a proxy |
178 | 192 |
|
179 | | - def reset_post_sim_init(self) -> None: # pragma: no cover |
180 | | - """Perform any once-per-episode setup.""" |
181 | | - pass |
182 | | - |
183 | 193 | @abstractmethod # pragma: no cover |
184 | 194 | def get_obs(self) -> Any: |
185 | 195 | """Return the observation.""" |
@@ -481,5 +491,30 @@ def get_obs(self): |
481 | 491 | ] |
482 | 492 |
|
483 | 493 |
|
| 494 | +class ResourceRewardWeight(Observation): |
| 495 | + """Observation for the weight of any :class:`ResourceReward`.""" |
| 496 | + |
| 497 | + def __init__( |
| 498 | + self, name="resource_reward_weight", norm: Union[float, np.ndarray] = 1.0 |
| 499 | + ): |
| 500 | + """Include the resource reward weight in the observation. |
| 501 | +
|
| 502 | + Args: |
| 503 | + name: Name of the observation. |
| 504 | + norm: Value to normalize the resource reward weight by. If a vector, it should |
| 505 | + be the same length as the number of resource rewards in the environment. |
| 506 | + """ |
| 507 | + super().__init__(name=name) |
| 508 | + self.norm = norm |
| 509 | + |
| 510 | + def reset_overwrite_previous(self) -> None: # pragma: no cover |
| 511 | + """Prepare weights to be saved by rewarders.""" |
| 512 | + self.weight_vector = [] |
| 513 | + |
| 514 | + def get_obs(self) -> float: |
| 515 | + """Return the resource reward weight.""" |
| 516 | + return np.array(self.weight_vector) / self.norm |
| 517 | + |
| 518 | + |
484 | 519 | __doc_title__ = "Backend" |
485 | 520 | __all__ = ["ObservationBuilder"] |
0 commit comments