Skip to content

Commit 6fac409

Browse files
author
EnliteAI Bot
committed
Refactor logger injection in rollout runners to improve logging flexibility and enable custom logger usage.
(Issue - null)
1 parent 54f659d commit 6fac409

2 files changed

Lines changed: 9 additions & 4 deletions

File tree

maze/core/rollout/rollout_runner.py

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,6 @@
22

33
from __future__ import annotations
44

5-
import logging
65
import os
76
import time
87
import traceback
@@ -27,9 +26,6 @@
2726
import numpy as np
2827
from omegaconf import DictConfig, OmegaConf
2928

30-
logger = logging.getLogger('SEQUENTIAL RUNNER')
31-
logger.setLevel(logging.INFO)
32-
3329

3430
class RolloutRunner(Runner, ABC):
3531
"""General abstract class for rollout runners.
@@ -241,6 +237,7 @@ def run_interaction_loop(
241237
env_seeds: list[Any],
242238
agent_seeds: list[Any],
243239
deterministic: bool,
240+
logger: Logger,
244241
render: bool = False,
245242
after_reset_callback: Callable = None,
246243
) -> None:
@@ -252,6 +249,8 @@ def run_interaction_loop(
252249
:param n_episodes: Count of episodes to perform.
253250
:param env_seeds: The env seeds to be used for each episode.
254251
:param agent_seeds: The agent seeds to be used for each episode.
252+
:param deterministic: Argmax policy.
253+
:param logger: Logger instance for logging rollout events.
255254
:param render: Whether to render the environment after every step.
256255
:param after_reset_callback: If supplied, this will be executed after each episode to notify the observer.
257256
"""

maze/core/rollout/sequential_rollout_runner.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,8 @@
22

33
from __future__ import annotations
44

5+
import logging
6+
57
from maze.core.annotations import override
68
from maze.core.log_events.log_events_writer_registry import LogEventsWriterRegistry
79
from maze.core.log_events.log_events_writer_tsv import LogEventsWriterTSV
@@ -21,6 +23,9 @@
2123

2224
from tqdm import tqdm
2325

26+
logger = logging.getLogger('SEQUENTIAL RUNNER')
27+
logger.setLevel(logging.INFO)
28+
2429

2530
class SequentialRolloutRunner(RolloutRunner):
2631
"""Runs rollout in the local process. Useful for short rollouts or debugging.
@@ -107,6 +112,7 @@ def run_with(self, env: ConfigType, wrappers: CollectionOfConfigType, agent: Con
107112
env_seeds=env_seeds,
108113
agent_seeds=agent_seeds,
109114
deterministic=self.deterministic,
115+
logger=logger,
110116
)
111117
self.progress_bar.close()
112118
env.write_epoch_stats()

0 commit comments

Comments
 (0)