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
|
||||
|
||||
# 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),
|
||||
)
|
||||
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}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Reference in a new issue