feat(message passing): fixed tracing bug, todo: dense layer input dimension incorrect
This commit is contained in:
parent
027b13b30f
commit
3cb3c765cc
1 changed files with 5 additions and 3 deletions
|
|
@ -134,8 +134,8 @@ def _normalize_obs(obs, mean, var, eps=1e-8):
|
||||||
|
|
||||||
|
|
||||||
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm: jnp.ndarray):
|
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, segments_per_arm: jnp.ndarray):
|
||||||
num_segments = segments_per_arm.sum()
|
num_segments = int(segments_per_arm.sum())
|
||||||
num_arms = jnp.where(segments_per_arm > 0, 1, 0).sum()
|
num_arms = int(jnp.where(segments_per_arm > 0, 1, 0).sum())
|
||||||
|
|
||||||
@logged_jit
|
@logged_jit
|
||||||
def _filter_and_flatten(o) -> jnp.ndarray:
|
def _filter_and_flatten(o) -> jnp.ndarray:
|
||||||
|
|
@ -668,9 +668,11 @@ class PPOTrainer:
|
||||||
|
|
||||||
message_passer_params = {}
|
message_passer_params = {}
|
||||||
if self.morph_mode != MorphMode.CENTRALIZED:
|
if self.morph_mode != MorphMode.CENTRALIZED:
|
||||||
|
assert self.message_passer is not None, "MessagePasser is None"
|
||||||
|
|
||||||
message_passer_params = jax.vmap(
|
message_passer_params = jax.vmap(
|
||||||
lambda k: self.message_passer.init(
|
lambda k: self.message_passer.init(
|
||||||
k, self.sensor.apply(single_sensor_param, sample_obs)
|
k, self.sensor.apply(single_sensor_param, sample_obs), self.adj
|
||||||
)
|
)
|
||||||
)(message_passer_keys)
|
)(message_passer_keys)
|
||||||
self.logger.info(
|
self.logger.info(
|
||||||
|
|
|
||||||
Reference in a new issue