Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
.venv/
**/__pycache__/
ppo_gulf*/
*logs*/
models/
42 changes: 42 additions & 0 deletions Env/Callbacks.py
Original file line number Diff line number Diff line change
Expand Up @@ -192,3 +192,45 @@ def _on_step(self) -> bool:
def _log_if_present(self, info_dict, key, tb_key):
if key in info_dict:
self.logger.record(tb_key, info_dict[key])



class EpisodeStatsCallback(BaseCallback):
"""
Reproduce SB3's ep_rew_mean logic, but using info["epi"].
Aggregates over episodes finished during *this rollout* only.
"""

def __init__(self, prefix="rollout", verbose=0):
super().__init__(verbose)
self.prefix = prefix
self._epis_this_rollout = []

def _on_rollout_start(self) -> None:
# Reset at start of rollout
self._epis_this_rollout = []

def _on_step(self) -> bool:
dones = self.locals["dones"]
infos = self.locals["infos"]
for done, info in zip(dones, infos):
if done and "epi" in info and isinstance(info["epi"], dict):
# store one dict per finished episode
self._epis_this_rollout.append(info["epi"].copy())
return True

def _on_rollout_end(self) -> None:
print(self._epis_this_rollout)
if not self._epis_this_rollout:
return

# Aggregate like SB3: mean over just-this-rollout episodes
keys = set().union(*[epi.keys() for epi in self._epis_this_rollout])
for key in sorted(keys):
values = [epi[key] for epi in self._epis_this_rollout if key in epi]
# Only log numeric values
values = [float(v) for v in values if isinstance(v, (int, float))]
if values:
mean_val = sum(values) / len(values)
self.logger.record(f"{self.prefix}/{key}_mean", mean_val)

98 changes: 61 additions & 37 deletions Env/MariNav.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,7 @@ def __init__(
pairs: list[tuple[str, str]],
h3_resolution: int = H3_RESOLUTION,
wind_threshold: float = DEFAULT_WIND_THRESHOLD,
no_positive_rews: bool = False,
render_mode: str = None,
):
"""
Expand All @@ -98,6 +99,7 @@ def __init__(
self.wind_threshold = wind_threshold
self.render_mode = render_mode
self.prioritize_until = 0
self.no_positive_rews = no_positive_rews

self.pairs = pairs

Expand Down Expand Up @@ -291,6 +293,7 @@ def reset(
self.prev_h3 = None
self.step_count = 0
self.episode_reward = 0.0
self.episode_neg_reward = 0.0
self.episode_wind_penalty = 0.0
self.episode_fuel_penalty = 0.0
self.episode_eta_penalty = 0.0
Expand Down Expand Up @@ -332,14 +335,10 @@ def _calculate_rewards(
wind_speed: float,
move_direction: float,
wind_direction: float,
override_reward: bool,
) -> float:
"""
Calculates the various reward components for the current step.
"""
if override_reward:
return INVALID_MOVE_PENALTY

# Progress reward: based on reduction in distance to goal
progress_reward = PROGRESS_REWARD_FACTOR * (distance_before - distance_after)

Expand Down Expand Up @@ -408,6 +407,34 @@ def _get_current_wind_conditions(self) -> tuple[float, float, float]:
wind_direction = np.arctan2(wind_v, wind_u) # Radians

return wind_u, wind_v, wind_speed, wind_direction

def _finalize_epi(self, info: dict) -> dict:
"""
Attach episode statistics to info when an episode ends.
Mirrors how Monitor provides info["episode"].
"""
info.update(
{
"epi": {
"r": (
self.episode_neg_reward
if self.no_positive_rews
else self.episode_reward
),
"l": self.step_count,
"progress_reward": self.episode_progress_reward,
"frequency_reward": self.episode_frequency_reward,
"wind_penalty": self.episode_wind_penalty,
"fuel_penalty": self.episode_fuel_penalty,
"eta_penalty": self.episode_eta_penalty,
"base_step_penalty": self.episode_base_step_penalty,
"episode_reward": self.episode_reward,
"episode_neg_reward": self.episode_neg_reward,
}
}
)
return info


def step(
self, action: tuple[int, int]
Expand All @@ -427,9 +454,9 @@ def step(
- info (dict): A dictionary containing additional information about the step.
"""
self.total_steps += 1
self.step_count += 1
info: dict = {}
info["revisiting_loop"] = 0
override_reward = False

neighbor_idx, speed_level = action
# neighbor_idx, speed_level = action
Expand All @@ -443,48 +470,44 @@ def step(
else:
ring, neighbors = self.get_all_neighbors()

if not neighbors:
return (
self._get_observation(),
-1900,
True,
False,
{"reason": "no valid neighbors"},
)

# Step 1: Check ring index range
if neighbor_idx >= len(ring):
raise Exception("neighbor_idx out of ring bounds")

if not neighbors:
self.episode_reward += INVALID_MOVE_PENALTY
info = {"reason": "no valid neighbors"}
info = self._finalize_epi(info)
return (
self._get_observation(),
-1900,
INVALID_MOVE_PENALTY,
True,
False,
{"reason": "neighbor_idx out of ring bounds"},
info,
)

# Step 2: Get candidate neighbor from ring
candidate = ring[neighbor_idx]
# Step 3: Make sure it's in valid neighbors
if candidate not in neighbors:
self.episode_reward += INVALID_MOVE_PENALTY
info = {"reason": "selected ring neighbor not in valid neighbors"}
info = self._finalize_epi(info)
return (
self._get_observation(),
-1900,
INVALID_MOVE_PENALTY,
True,
False,
{"reason": "selected ring neighbor not in valid neighbors"},
info,
)

# Step 4: It's valid
selected_neighbor = candidate

if selected_neighbor == self.current_h3:
override_reward = True

try:
distance_before = self.shortest_path_length(self.current_h3, self.goal_h3)
self.prev_h3 = self.current_h3
self.current_h3 = selected_neighbor
self.step_count += 1
distance_after = self.shortest_path_length(selected_neighbor, self.goal_h3)
except nx.NetworkXNoPath:
# No valid path exists — strongly penalize the move
Expand Down Expand Up @@ -530,7 +553,6 @@ def step(
wind_speed,
move_direction,
wind_direction,
override_reward,
)

step_wind_penalty = calculated_rewards["wind_penalty"]
Expand All @@ -540,19 +562,18 @@ def step(
step_progress_reward = calculated_rewards["progress_reward"]
step_frequency_reward = calculated_rewards["frequency_reward"]
step_total_reward = calculated_rewards["total_reward"]
step_total_neg_reward = (
step_total_reward - step_progress_reward - step_frequency_reward
)

if terminated:
key = (self.start_h3, self.goal_h3)
self.visited_path_counts[key] = self.visited_path_counts.get(key, 0) + 1
print(
f"Hurray! Goal reached! Start_H3: {self.start_h3}, Goal_H3: {self.goal_h3}, "
f"Step Count: {self.step_count}, Episode Reward: {self.episode_reward}"
)
step_total_reward += GOAL_REWARD

step_total_reward = (
step_total_reward / self.max_distance
) * self.max_distance_reference
distance_scaler = self.max_distance_reference / self.max_distance
step_total_reward = step_total_reward * distance_scaler
step_total_neg_reward = step_total_neg_reward * distance_scaler

self.episode_reward += step_total_reward
self.episode_wind_penalty += step_wind_penalty
Expand All @@ -561,22 +582,27 @@ def step(
self.episode_base_step_penalty += step_base_step_penalty
self.episode_progress_reward += step_progress_reward
self.episode_frequency_reward += step_frequency_reward
self.true_reward_with_no_progress_reward = (
self.episode_reward - self.episode_progress_reward
)
self.episode_neg_reward += step_total_neg_reward

if done:
print("Done")
info.update(
{
"epi": {
"r": self.episode_reward,
"r": (
self.episode_neg_reward
if self.no_positive_rews
else self.episode_reward
),
"l": self.step_count,
"progress_reward": self.episode_progress_reward,
"frequency_reward": self.episode_frequency_reward,
"wind_penalty": self.episode_wind_penalty,
"fuel_penalty": self.episode_fuel_penalty,
"eta_penalty": self.episode_eta_penalty,
"base_step_penalty": self.episode_base_step_penalty,
"episode_reward": self.episode_reward,
"episode_neg_reward": self.episode_neg_reward,
},
"visited_path_counts": {
f"{s}->{g}": count
Expand All @@ -586,7 +612,6 @@ def step(
f"{s}->{g}": count
for (s, g), count in self.pair_selection_counts.items()
},
"true_etr": self.true_reward_with_no_progress_reward,
}
)

Expand All @@ -601,7 +626,6 @@ def step(
"current_h3": self.current_h3,
"prev_h3": self.prev_h3,
"distance_to_goal": distance_after,
"override_reward": override_reward,
# Step-level reward components with self_ prefix
"self_progress_reward": step_progress_reward,
"self_frequency_reward": step_frequency_reward,
Expand All @@ -624,7 +648,7 @@ def step(

return (
self._get_observation(speed, wind_direction),
step_total_reward,
step_total_neg_reward if self.no_positive_rews else step_total_reward,
terminated,
truncated,
info,
Expand Down
3 changes: 2 additions & 1 deletion train_PPO.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,9 +137,10 @@ def _init():

step_logger = StepRewardLoggerCallback()
info_logging_callback = InfoLoggingCallback()
ep_stat_callback = EpisodeStatsCallback()

callback = CallbackList(
[eval_callback, early_stop, step_logger, info_logging_callback]
[eval_callback, early_stop, step_logger, info_logging_callback, ep_stat_callback]
)

# Train the model
Expand Down
Loading
Loading