1
Fork 0

feat(message passing): used config value for message passing repetitions

This commit is contained in:
Robin Meersman 2026-05-02 23:15:55 +02:00
parent 6114e36e27
commit 027b13b30f
4 changed files with 14 additions and 3 deletions

View file

@ -29,4 +29,4 @@ critic:
message_passing_steps: 4
# Connectivity topology (e.g., ring, fully_connected)
topology_type: "ring"
topology_type: "fully_connected"

View file

@ -0,0 +1,6 @@
# 2 Arms Morphology Configuration
segments_per_arm: [4, 0, 4, 0, 0]
use_p_control: true
use_torque_control: false
morph_mode: FULLY_CONNECTED

View file

@ -42,7 +42,7 @@ def main(dict_cfg: DictConfig):
base_dir=os.path.dirname(run_dir),
)
logger = get_logger()
logger.set_level(logging.DEBUG)
logger.set_level(logging.INFO)
logger.info(f"Hydra-initialized run: {run_name}")
logger.info(f"Output directory: {run_dir}")

View file

@ -221,10 +221,15 @@ def _get_action_and_value_noise(
adj_matrix: jnp.ndarray,
):
def apply_message_passing(h):
assert message_passer is not None, (
"message passing shouldn't be performed if morphology mode is Centralized"
)
return message_passer.apply(agent_state.params["message_passer_params"], h, adj_matrix)
hidden = apply_per_node(sensor, agent_state.params["sensor_params"], next_obs)
hidden = jax.vmap(apply_message_passing)(hidden)
if message_passer is not None:
hidden = jax.vmap(apply_message_passing)(hidden)
hidden_critic = apply_shared(
feature_extractor, agent_state.params["feature_extractor_params"], next_obs