feat(message passing): used config value for message passing repetitions
This commit is contained in:
parent
6114e36e27
commit
027b13b30f
4 changed files with 14 additions and 3 deletions
|
|
@ -29,4 +29,4 @@ critic:
|
||||||
message_passing_steps: 4
|
message_passing_steps: 4
|
||||||
|
|
||||||
# Connectivity topology (e.g., ring, fully_connected)
|
# Connectivity topology (e.g., ring, fully_connected)
|
||||||
topology_type: "ring"
|
topology_type: "fully_connected"
|
||||||
|
|
|
||||||
6
configs/morphology/2_arms_decentralized.yaml
Normal file
6
configs/morphology/2_arms_decentralized.yaml
Normal 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
|
||||||
|
|
@ -42,7 +42,7 @@ def main(dict_cfg: DictConfig):
|
||||||
base_dir=os.path.dirname(run_dir),
|
base_dir=os.path.dirname(run_dir),
|
||||||
)
|
)
|
||||||
logger = get_logger()
|
logger = get_logger()
|
||||||
logger.set_level(logging.DEBUG)
|
logger.set_level(logging.INFO)
|
||||||
logger.info(f"Hydra-initialized run: {run_name}")
|
logger.info(f"Hydra-initialized run: {run_name}")
|
||||||
logger.info(f"Output directory: {run_dir}")
|
logger.info(f"Output directory: {run_dir}")
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -221,10 +221,15 @@ def _get_action_and_value_noise(
|
||||||
adj_matrix: jnp.ndarray,
|
adj_matrix: jnp.ndarray,
|
||||||
):
|
):
|
||||||
def apply_message_passing(h):
|
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)
|
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 = 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(
|
hidden_critic = apply_shared(
|
||||||
feature_extractor, agent_state.params["feature_extractor_params"], next_obs
|
feature_extractor, agent_state.params["feature_extractor_params"], next_obs
|
||||||
|
|
|
||||||
Reference in a new issue