fix(ppo trainer): added check if message passing is None
This commit is contained in:
parent
23fcee2110
commit
7f5b730cdf
3 changed files with 15 additions and 2 deletions
7
configs/morphology/5_arms_full_fullconnected.yaml
Normal file
7
configs/morphology/5_arms_full_fullconnected.yaml
Normal file
|
|
@ -0,0 +1,7 @@
|
||||||
|
# 5 Arms Full Morphology Configuration
|
||||||
|
# Baseline 5-arm brittle star.
|
||||||
|
|
||||||
|
segments_per_arm: [4, 4, 4, 4, 4]
|
||||||
|
use_p_control: true
|
||||||
|
use_torque_control: false
|
||||||
|
morph_mode: FULLY_CONNECTED
|
||||||
|
|
@ -119,7 +119,8 @@ def get_action_and_value(
|
||||||
):
|
):
|
||||||
hidden_sensor = sensor_apply(params["sensor_params"], x)
|
hidden_sensor = sensor_apply(params["sensor_params"], x)
|
||||||
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
|
hidden_critic = feature_extractor_apply(params["feature_extractor_params"], x)
|
||||||
hidden_sensor = message_passer(params["message_passer_params"], hidden_sensor)
|
if message_passer is not None:
|
||||||
|
hidden_sensor = message_passer(params["message_passer_params"], hidden_sensor)
|
||||||
|
|
||||||
debug.callback(logger.debug, f"[SHAPE] hidden_sensor: {hidden_sensor.shape}")
|
debug.callback(logger.debug, f"[SHAPE] hidden_sensor: {hidden_sensor.shape}")
|
||||||
debug.callback(logger.debug, f"[SHAPE] hidden_critic: {hidden_critic.shape}")
|
debug.callback(logger.debug, f"[SHAPE] hidden_critic: {hidden_critic.shape}")
|
||||||
|
|
|
||||||
|
|
@ -514,7 +514,12 @@ class PPOTrainer:
|
||||||
return jax.vmap(lambda x_in: self.message_passer.apply(p, x_in))(x)
|
return jax.vmap(lambda x_in: self.message_passer.apply(p, x_in))(x)
|
||||||
|
|
||||||
self._ppo = PPO(
|
self._ppo = PPO(
|
||||||
self.ppo, apply_sensor, apply_actor, apply_critic, apply_feature, apply_message_passer
|
self.ppo,
|
||||||
|
apply_sensor,
|
||||||
|
apply_actor,
|
||||||
|
apply_critic,
|
||||||
|
apply_feature,
|
||||||
|
apply_message_passer if self.message_passer is not None else None,
|
||||||
)
|
)
|
||||||
|
|
||||||
self.agent_state = self._init_agent_state()
|
self.agent_state = self._init_agent_state()
|
||||||
|
|
|
||||||
Reference in a new issue