From 517c3134ca0eb68478b20d3e6765907f2476d8b5 Mon Sep 17 00:00:00 2001 From: Robin Meersman Date: Thu, 7 May 2026 14:00:39 +0200 Subject: [PATCH] fix: rewrote info logger calls to debug calls for debug related logs --- .../environment/obs_processing.py | 16 ++++++++-------- src/brittle_star_project/ppo.py | 2 ++ 2 files changed, 10 insertions(+), 8 deletions(-) diff --git a/src/brittle_star_project/environment/obs_processing.py b/src/brittle_star_project/environment/obs_processing.py index a57f00b..41afe57 100644 --- a/src/brittle_star_project/environment/obs_processing.py +++ b/src/brittle_star_project/environment/obs_processing.py @@ -128,18 +128,18 @@ def create_obs_processor( for k, v in obs.items(): if hasattr(v, "shape"): size = v.size - logger11.info(f"[RAW] {k}: shape={v.shape}, size={size}") + logger11.debug(f"[RAW] {k}: shape={v.shape}, size={size}") total += size else: - logger11.info(f"[RAW] {k}: non-array") + logger11.debug(f"[RAW] {k}: non-array") - logger11.info(f"[RAW TOTAL FEATURES]: {total}") + logger11.debug(f"[RAW TOTAL FEATURES]: {total}") output = {} num_agents = needed_copies # IMPORTANT: number of MLPs for key, arr in obs.items(): if key not in ordered_keys or arr.size == 0: continue - logger11.info(f"[INPUT] {key}: {arr.shape}") + logger11.debug(f"[INPUT] {key}: {arr.shape}") if arr.ndim == 0: arr = arr.reshape(1) # -------- CENTRALIZED -------- @@ -154,7 +154,7 @@ def create_obs_processor( for i, agent_id in enumerate(agent_indices): idx = segment_indices[i] taken = jnp.take(arr, idx, axis=0) # (segs, ...) - logger11.info(f"WHY {taken.shape}") + logger11.debug(f"WHY {taken.shape}") # pad to 4 pad_len = 4 - taken.shape[0] padded = jnp.pad(taken, [(0, pad_len)] + [(0, 0)] * (taken.ndim - 1)) @@ -182,7 +182,7 @@ def create_obs_processor( else: out = jnp.repeat(arr[None, :], num_agents, axis=0) - logger11.info(f"[OUTPUT] {key}: {out.shape}") + logger11.debug(f"[OUTPUT] {key}: {out.shape}") output[key] = out return output @@ -220,8 +220,8 @@ def create_obs_processor( processed = _normalize_features(processed) processed = _split_to_agents(processed, morph_mode) flat = _flatten_features(processed) # (num_arms, total_feat) - logger11.info(f"[FLATTENED FINAL] shape: {flat.shape}") - logger11.info(f"[PER AGENT] example row 0 shape: {flat[0].shape}") + logger11.debug(f"[FLATTENED FINAL] shape: {flat.shape}") + logger11.debug(f"[PER AGENT] example row 0 shape: {flat[0].shape}") return _flatten_features(processed) # (agents, feat) return jax.jit(jax.vmap(_process_single)) diff --git a/src/brittle_star_project/ppo.py b/src/brittle_star_project/ppo.py index 79fea1c..fffc26b 100644 --- a/src/brittle_star_project/ppo.py +++ b/src/brittle_star_project/ppo.py @@ -119,6 +119,8 @@ def get_action_and_value( ): hidden_sensor = sensor_apply(params["sensor_params"], x) hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x) + + # only apply message passing in decentralized context if message_passer is not None: hidden_sensor = message_passer(params["message_passer_params"], hidden_sensor)