Skip to content

Commit 645b335

Browse files
DavidePagliericopybara-github
authored andcommitted
exposing verbose options
PiperOrigin-RevId: 799570510 Change-Id: Ia60b44ba993ab5347fc68fa01a294d84dba796bd
1 parent b83e1d4 commit 645b335

2 files changed

Lines changed: 20 additions & 6 deletions

File tree

concordia/environment/engines/parallel.py

Lines changed: 11 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -150,7 +150,10 @@ def terminate(
150150
return should_terminate_string == entity_lib.BINARY_OPTIONS['affirmative']
151151

152152
def make_observation(
153-
self, game_master: entity_lib.Entity, entity: entity_lib.Entity
153+
self,
154+
game_master: entity_lib.Entity,
155+
entity: entity_lib.Entity,
156+
verbose: bool = False,
154157
) -> str:
155158
"""Make an observation for a game object."""
156159
observation = game_master.act(
@@ -161,7 +164,12 @@ def make_observation(
161164
output_type=entity_lib.OutputType.MAKE_OBSERVATION,
162165
)
163166
)
164-
print(f'Observation: {observation} for {entity.name}')
167+
if verbose:
168+
print(
169+
termcolor.colored(
170+
f'Observation: {observation} for {entity.name}', _PRINT_COLOR
171+
)
172+
)
165173
return observation
166174

167175
@override
@@ -200,7 +208,7 @@ def run_loop(
200208
tasks = {}
201209
for entity in entities:
202210
tasks[entity.name] = functools.partial(
203-
entity.observe, self.make_observation(game_master, entity)
211+
entity.observe, self.make_observation(game_master, entity, verbose)
204212
)
205213
concurrency.run_tasks(tasks, executor=executor)
206214

concordia/prefabs/simulation/questionnaire_simulation.py

Lines changed: 9 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@ def __init__(
4848
embedder: Callable[[str], np.ndarray],
4949
engine: parallel.ParallelQuestionnaireEngine | None = None,
5050
max_workers: int | None = None,
51+
verbose: bool = False,
5152
):
5253
"""Initialize the simulation object.
5354
@@ -68,10 +69,12 @@ def __init__(
6869
parallel.ParallelQuestionnaireEngine.
6970
max_workers: the maximum number of workers to use in the engine's
7071
ThreadPoolExecutor, if the default engine is used.
72+
verbose: Whether to print verbose output.
7173
"""
7274
self._config = config
7375
self._model = model
7476
self._embedder = embedder
77+
self._verbose = verbose
7578
if engine is None:
7679
if not max_workers:
7780
self._engine = parallel.ParallelQuestionnaireEngine()
@@ -193,11 +196,13 @@ def add_entity(
193196
# Check if a pre-loaded memory state was passed in the entity's params.
194197
memory_state = instance_config.params.get("memory_state")
195198
if memory_state:
196-
print(f"Found pre-loaded memory state for {entity.name}. Setting it.")
199+
if self._verbose:
200+
print(f"Found pre-loaded memory state for {entity.name}. Setting it.")
197201
try:
198202
memory_component = entity.get_component("__memory__")
199203
memory_component.set_state(memory_state)
200-
print(f"Successfully set pre-loaded memories for {entity.name}.")
204+
if self._verbose:
205+
print(f"Successfully set pre-loaded memories for {entity.name}.")
201206
except (KeyError, TypeError, ValueError) as e:
202207
print(f"Error setting pre-loaded memory for {entity.name}: {e}")
203208

@@ -372,7 +377,8 @@ def save_checkpoint(self, step: int, checkpoint_path: str):
372377
try:
373378
with open(checkpoint_file, "w") as f:
374379
json.dump(checkpoint_data, f, indent=2)
375-
print(f"Step {step}: Saved checkpoint to {checkpoint_file}")
380+
if self._verbose:
381+
print(f"Step {step}: Saved checkpoint to {checkpoint_file}")
376382
except IOError as e:
377383
print(f"Error saving checkpoint at step {step}: {e}")
378384

0 commit comments

Comments
 (0)