merge: merged dev into branch
This commit is contained in:
parent
61064ca70e
commit
43de6a9670
3 changed files with 79 additions and 153 deletions
|
|
@ -48,7 +48,8 @@ class MessagePasser(nn.Module):
|
||||||
messages = nn.Dense(self.hidden_dim)(x)
|
messages = nn.Dense(self.hidden_dim)(x)
|
||||||
messages = nn.tanh(messages)
|
messages = nn.tanh(messages)
|
||||||
|
|
||||||
agg = adj_matrix # if mean is wanted: adj_matrix / (adj.sum(axis=-1, keepdims=True) + 1e-8)
|
# note: if mean is wanted: adj_matrix / (adj.sum(axis=-1, keepdims=True) + 1e-8)
|
||||||
|
agg = adj_matrix
|
||||||
aggregated = agg @ messages
|
aggregated = agg @ messages
|
||||||
|
|
||||||
x_concat = jnp.concatenate([x, aggregated], axis=-1)
|
x_concat = jnp.concatenate([x, aggregated], axis=-1)
|
||||||
|
|
|
||||||
|
|
@ -2,6 +2,9 @@ import jax
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
from typing import Dict, Tuple, Optional
|
from typing import Dict, Tuple, Optional
|
||||||
|
|
||||||
|
from brittle_star_project.environment.env_config import MorphMode
|
||||||
|
from experiment_logger import get_logger
|
||||||
|
|
||||||
_JOINT_SCALED_KEYS = frozenset(
|
_JOINT_SCALED_KEYS = frozenset(
|
||||||
{
|
{
|
||||||
"joint_position",
|
"joint_position",
|
||||||
|
|
@ -19,7 +22,11 @@ _SEGMENT_SCALED_KEYS = frozenset(
|
||||||
|
|
||||||
|
|
||||||
def create_obs_processor(
|
def create_obs_processor(
|
||||||
bounds_dict: Dict[str, Tuple[float, float]], padding_masks: Optional[Dict] = None
|
bounds_dict: Dict[str, Tuple[float, float]],
|
||||||
|
num_segments: int,
|
||||||
|
num_arms: int,
|
||||||
|
padding_masks: Optional[Dict] = None,
|
||||||
|
morph_mode: MorphMode = MorphMode.CENTRALIZED,
|
||||||
):
|
):
|
||||||
def _add_derived_features(obs: dict) -> dict:
|
def _add_derived_features(obs: dict) -> dict:
|
||||||
new_obs = dict(obs)
|
new_obs = dict(obs)
|
||||||
|
|
@ -52,6 +59,8 @@ def create_obs_processor(
|
||||||
return normalized
|
return normalized
|
||||||
|
|
||||||
def _pad_features(obs: dict) -> dict:
|
def _pad_features(obs: dict) -> dict:
|
||||||
|
assert padding_masks is not None
|
||||||
|
|
||||||
padded = {}
|
padded = {}
|
||||||
for key, arr in obs.items():
|
for key, arr in obs.items():
|
||||||
if key in _JOINT_SCALED_KEYS:
|
if key in _JOINT_SCALED_KEYS:
|
||||||
|
|
@ -73,13 +82,55 @@ def create_obs_processor(
|
||||||
"robot_direction_to_target",
|
"robot_direction_to_target",
|
||||||
"segment_contact",
|
"segment_contact",
|
||||||
]
|
]
|
||||||
|
|
||||||
values = []
|
values = []
|
||||||
for key in ordered_keys:
|
|
||||||
if key in obs:
|
for key in sorted(obs.keys()):
|
||||||
arr = jnp.asarray(obs[key]).flatten()
|
if key not in ordered_keys:
|
||||||
if arr.size > 0:
|
continue
|
||||||
values.append(arr)
|
v = obs[key]
|
||||||
return jnp.concatenate(values)
|
|
||||||
|
# skip empty arrays and scalars
|
||||||
|
if v.size == 0:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# reshape scalars
|
||||||
|
if v.ndim == 0:
|
||||||
|
v = v.reshape(1)
|
||||||
|
|
||||||
|
# -------- CENTRALIZED --------
|
||||||
|
if morph_mode == MorphMode.CENTRALIZED:
|
||||||
|
values.append(v.reshape(1, -1)) # (1, feat)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# -------- SPLIT TO SEGMENTS --------
|
||||||
|
if key in _JOINT_SCALED_KEYS:
|
||||||
|
if morph_mode == MorphMode.SEGMENT:
|
||||||
|
center_size = num_arms * 3 * 2
|
||||||
|
v_center = v[:center_size].reshape(num_arms, 3 * 2) # (arms, 6)
|
||||||
|
v_segs = v[center_size:].reshape(-1, 2) # (segs, 2)
|
||||||
|
values.append(jnp.concatenate([v_center, v_segs], axis=0)) # (arms+segs, ?)
|
||||||
|
continue
|
||||||
|
v = v.reshape(num_arms, -1) # (n_arms, 2)
|
||||||
|
|
||||||
|
elif key in _SEGMENT_SCALED_KEYS:
|
||||||
|
v = v[:, None] # (segments, 1)
|
||||||
|
|
||||||
|
else:
|
||||||
|
# global key, broadcast to all nodes
|
||||||
|
n_nodes = (num_segments + num_arms) if morph_mode == MorphMode.SEGMENT else num_arms
|
||||||
|
v = jnp.repeat(v[None, :], n_nodes, axis=0) # (n_nodes, feat)
|
||||||
|
|
||||||
|
# -------- SEGMENT MODE --------
|
||||||
|
if morph_mode == MorphMode.SEGMENT:
|
||||||
|
values.append(v) # (n_nodes, feat)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# -------- ARM MODE --------
|
||||||
|
v = v.reshape(num_arms, -1)
|
||||||
|
values.append(v) # (n_arms, feat)
|
||||||
|
|
||||||
|
return jnp.concatenate(values, axis=-1)
|
||||||
|
|
||||||
def _process_single(obs_dict: dict) -> jnp.ndarray:
|
def _process_single(obs_dict: dict) -> jnp.ndarray:
|
||||||
processed = _add_derived_features(obs_dict)
|
processed = _add_derived_features(obs_dict)
|
||||||
|
|
|
||||||
|
|
@ -32,18 +32,7 @@ from brittle_star_project.utils import logged_jit
|
||||||
|
|
||||||
logger11 = get_logger()
|
logger11 = get_logger()
|
||||||
# TODO: move to config
|
# TODO: move to config
|
||||||
_ALLOWED_OBS_KEYS = {
|
|
||||||
"joint_position",
|
|
||||||
"joint_velocity",
|
|
||||||
"joint_actuator_force",
|
|
||||||
"actuator_force",
|
|
||||||
"disk_position",
|
|
||||||
"disk_rotation",
|
|
||||||
"disk_linear_velocity",
|
|
||||||
"disk_angular_velocity",
|
|
||||||
"unit_xy_direction_to_target",
|
|
||||||
"xy_distance_to_target",
|
|
||||||
}
|
|
||||||
# TODO: clip scaled reward?
|
# TODO: clip scaled reward?
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -129,81 +118,6 @@ def _linear_schedule(count, minibatch_count, update_epochs, num_iterations, lear
|
||||||
return learning_rate * frac
|
return learning_rate * frac
|
||||||
|
|
||||||
|
|
||||||
@logged_jit
|
|
||||||
def _normalize_obs(obs, mean, var, eps=1e-8):
|
|
||||||
return jnp.clip((obs - mean) / jnp.sqrt(var + eps), -10.0, 10.0)
|
|
||||||
|
|
||||||
# TODO: update to work with new obs_processor
|
|
||||||
def _convert_obs_dict_to_array_morphology(obs_dict, morph_mode, num_segments: int, num_arms: int):
|
|
||||||
@logged_jit
|
|
||||||
def _filter_and_flatten(o) -> jnp.ndarray:
|
|
||||||
# vmap feeds one env at a time — v has NO batch dim here
|
|
||||||
# shapes are e.g. (n_features,) or (n_nodes, feat)
|
|
||||||
values = []
|
|
||||||
|
|
||||||
for key in sorted(o.keys()):
|
|
||||||
if key not in _ALLOWED_OBS_KEYS:
|
|
||||||
continue
|
|
||||||
v = o[key]
|
|
||||||
if v.size == 0:
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- CENTRALIZED --------
|
|
||||||
if morph_mode == MorphMode.CENTRALIZED:
|
|
||||||
values.append(v.reshape(1, -1)) # (1, feat)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- SPLIT TO SEGMENTS --------
|
|
||||||
if key in _JOINT_SCALED_KEYS:
|
|
||||||
if morph_mode == MorphMode.SEGMENT:
|
|
||||||
center_size = num_arms * 3 * 2
|
|
||||||
v_center = v[:center_size].reshape(num_arms, 3 * 2) # (arms, 6)
|
|
||||||
v_segs = v[center_size:].reshape(-1, 2) # (segs, 2)
|
|
||||||
values.append(jnp.concatenate([v_center, v_segs], axis=0)) # (arms+segs, ?)
|
|
||||||
continue
|
|
||||||
v = v.reshape(num_arms, -1) # (n_arms, 2)
|
|
||||||
|
|
||||||
elif key in _SEGMENT_SCALED_KEYS:
|
|
||||||
v = v[:, None] # (segments, 1)
|
|
||||||
|
|
||||||
else:
|
|
||||||
# global key, broadcast to all nodes
|
|
||||||
n_nodes = (num_segments + num_arms) if morph_mode == MorphMode.SEGMENT else num_arms
|
|
||||||
v = jnp.repeat(v[None, :], n_nodes, axis=0) # (n_nodes, feat)
|
|
||||||
|
|
||||||
# -------- SEGMENT MODE --------
|
|
||||||
if morph_mode == MorphMode.SEGMENT:
|
|
||||||
values.append(v) # (n_nodes, feat)
|
|
||||||
continue
|
|
||||||
|
|
||||||
# -------- ARM MODE --------
|
|
||||||
v = v.reshape(num_arms, -1)
|
|
||||||
values.append(v) # (n_arms, feat)
|
|
||||||
|
|
||||||
return jnp.concatenate(values, axis=-1) # (n_nodes, total_feat)
|
|
||||||
|
|
||||||
return jax.vmap(_filter_and_flatten)(obs_dict)
|
|
||||||
# output: (batch, n_nodes, total_feat)
|
|
||||||
|
|
||||||
|
|
||||||
# Observation keys whose size scales with the number of joints (2 per segment).
|
|
||||||
_JOINT_SCALED_KEYS = frozenset(
|
|
||||||
{ # TODO CODE SMELL
|
|
||||||
"joint_position",
|
|
||||||
"joint_velocity",
|
|
||||||
"joint_actuator_force",
|
|
||||||
"actuator_force",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
# Observation keys whose size scales with the number of segments (1 per segment).
|
|
||||||
_SEGMENT_SCALED_KEYS = frozenset(
|
|
||||||
{
|
|
||||||
"segment_contact",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def _get_action_and_value_noise(
|
def _get_action_and_value_noise(
|
||||||
sensor: nn.Module,
|
sensor: nn.Module,
|
||||||
feature_extractor: nn.Module,
|
feature_extractor: nn.Module,
|
||||||
|
|
@ -249,7 +163,6 @@ def _get_action_and_value_noise(
|
||||||
return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
|
return flat_clipped_action, raw_action, logprob, value.squeeze(-1), mean, std, key
|
||||||
|
|
||||||
|
|
||||||
# TODO: update to work vectorized (sensor, actor, message passer) + message passing
|
|
||||||
def _step_once(
|
def _step_once(
|
||||||
carry,
|
carry,
|
||||||
_,
|
_,
|
||||||
|
|
@ -328,9 +241,7 @@ def _reward_fn(env_state, next_env_state):
|
||||||
return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty)
|
return jnp.where(next_env_state.terminated, 50.0, clipped_env_reward - penalty)
|
||||||
|
|
||||||
|
|
||||||
def _step_env_wrapped(
|
def _step_env_wrapped(episode_stats, env_state, action, env_step_fn, obs_processor):
|
||||||
episode_stats, env_state, action, env_step_fn, morph_mode, num_segments: int, num_arms: int, # TODO: obs_processor
|
|
||||||
):
|
|
||||||
next_env_state = env_step_fn(env_state, action)
|
next_env_state = env_step_fn(env_state, action)
|
||||||
|
|
||||||
reward = _reward_fn(env_state, next_env_state)
|
reward = _reward_fn(env_state, next_env_state)
|
||||||
|
|
@ -355,13 +266,10 @@ def _step_env_wrapped(
|
||||||
episode_stats,
|
episode_stats,
|
||||||
next_env_state,
|
next_env_state,
|
||||||
(
|
(
|
||||||
_convert_obs_dict_to_array_morphology(
|
obs_processor(next_env_state.observations),
|
||||||
next_env_state.observations, morph_mode, num_segments, num_arms
|
|
||||||
),
|
|
||||||
reward,
|
reward,
|
||||||
done,
|
done,
|
||||||
),
|
),
|
||||||
# TODO (obs_processor(next_env_state.observations), reward, done),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -502,7 +410,7 @@ class PPOTrainer:
|
||||||
self.key = jax.random.PRNGKey(self.experiment.seed)
|
self.key = jax.random.PRNGKey(self.experiment.seed)
|
||||||
|
|
||||||
self.morph_mode = self.cfg.morphology.morph_mode
|
self.morph_mode = self.cfg.morphology.morph_mode
|
||||||
|
|
||||||
self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32)
|
self.segments_per_arm = jnp.asarray(self.cfg.morphology.segments_per_arm, dtype=jnp.int32)
|
||||||
self.num_segments = self.segments_per_arm.sum().item()
|
self.num_segments = self.segments_per_arm.sum().item()
|
||||||
self.num_arms = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item()
|
self.num_arms = jnp.where(self.segments_per_arm > 0, 1, 0).sum().item()
|
||||||
|
|
@ -527,6 +435,9 @@ class PPOTrainer:
|
||||||
# Build the centralized observation processor: derive -> normalize -> pad -> flatten.
|
# Build the centralized observation processor: derive -> normalize -> pad -> flatten.
|
||||||
self.obs_processor = create_obs_processor(
|
self.obs_processor = create_obs_processor(
|
||||||
bounds_dict=self.cfg.obs_bounds.to_bounds_dict(),
|
bounds_dict=self.cfg.obs_bounds.to_bounds_dict(),
|
||||||
|
num_segments=self.num_segments,
|
||||||
|
num_arms=self.num_arms,
|
||||||
|
morph_mode=self.morph_mode,
|
||||||
padding_masks=self.env.padding_masks,
|
padding_masks=self.env.padding_masks,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
@ -540,10 +451,7 @@ class PPOTrainer:
|
||||||
step_env_fn=partial(
|
step_env_fn=partial(
|
||||||
_step_env_wrapped,
|
_step_env_wrapped,
|
||||||
env_step_fn=self.env.step,
|
env_step_fn=self.env.step,
|
||||||
morph_mode=self.morph_mode,
|
obs_processor=self.obs_processor,
|
||||||
num_segments=self.num_segments,
|
|
||||||
num_arms=self.num_arms,
|
|
||||||
# TODO: obs_processor=self.obs_processor,
|
|
||||||
),
|
),
|
||||||
sensor=self.sensor,
|
sensor=self.sensor,
|
||||||
feature_extractor=self.feature_extractor,
|
feature_extractor=self.feature_extractor,
|
||||||
|
|
@ -605,6 +513,7 @@ class PPOTrainer:
|
||||||
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
|
self.segments_per_arm.sum() + jnp.where(self.segments_per_arm > 0, 1, 0).sum()
|
||||||
).item()
|
).item()
|
||||||
|
|
||||||
|
# scale actor output with size of model --> more models ==> less actions needed per model
|
||||||
actor = Actor(action_dim=self.env.single_action_space.shape[0] // needed_copies)
|
actor = Actor(action_dim=self.env.single_action_space.shape[0] // needed_copies)
|
||||||
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
sensor = GenericDenseLayersWithActivation(layer_sizes=[300, 300, 300])
|
||||||
message_passer: Optional[nn.Module] = (
|
message_passer: Optional[nn.Module] = (
|
||||||
|
|
@ -628,15 +537,11 @@ class PPOTrainer:
|
||||||
)
|
)
|
||||||
|
|
||||||
dummy_reset = self.env.reset(seed=0)
|
dummy_reset = self.env.reset(seed=0)
|
||||||
|
|
||||||
for k, v in dummy_reset.observations.items():
|
for k, v in dummy_reset.observations.items():
|
||||||
self.logger.debug(k, v.shape)
|
self.logger.debug(k, v.shape)
|
||||||
sample_obs = _convert_obs_dict_to_array_morphology(
|
|
||||||
dummy_reset.observations,
|
sample_obs = self.obs_processor(dummy_reset.observations)[0] # take first env
|
||||||
self.morph_mode,
|
|
||||||
self.num_segments,
|
|
||||||
self.num_arms,
|
|
||||||
)[0] # take first env
|
|
||||||
|
|
||||||
self.logger.debug(f"[_init_agent_state] sample_obs: {sample_obs.shape}")
|
self.logger.debug(f"[_init_agent_state] sample_obs: {sample_obs.shape}")
|
||||||
self.obs_mean = jnp.zeros((sample_obs.shape[-1],))
|
self.obs_mean = jnp.zeros((sample_obs.shape[-1],))
|
||||||
|
|
@ -703,7 +608,6 @@ class PPOTrainer:
|
||||||
critic_params = self.critic.init(critic_key, critic_input)
|
critic_params = self.critic.init(critic_key, critic_input)
|
||||||
self.logger.debug(
|
self.logger.debug(
|
||||||
f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}"
|
f"[_init_agent_state] critic_params: {jax.tree.map(lambda x: x.shape, critic_params)}"
|
||||||
# TODO: sample_obs = self.obs_processor(dummy_reset.observations)[0] # take first env
|
|
||||||
)
|
)
|
||||||
|
|
||||||
return TrainState.create(
|
return TrainState.create(
|
||||||
|
|
@ -744,28 +648,6 @@ class PPOTrainer:
|
||||||
returned_episode_lengths=jnp.zeros(self.ppo.num_envs, dtype=jnp.int32),
|
returned_episode_lengths=jnp.zeros(self.ppo.num_envs, dtype=jnp.int32),
|
||||||
)
|
)
|
||||||
|
|
||||||
def _update_obs_stats(self, obs: jnp.ndarray):
|
|
||||||
batch_mean = jnp.mean(obs, axis=(0, 1))
|
|
||||||
batch_var = jnp.var(obs, axis=(0, 1))
|
|
||||||
batch_count = obs.shape[0]
|
|
||||||
|
|
||||||
self.logger.debug(f"Batch mean shape: {batch_mean.shape}")
|
|
||||||
self.logger.debug(f"obs_mean shape: {self.obs_mean.shape}")
|
|
||||||
|
|
||||||
delta = batch_mean - self.obs_mean
|
|
||||||
total_count = self.obs_count + batch_count
|
|
||||||
|
|
||||||
new_mean = self.obs_mean + delta * batch_count / total_count
|
|
||||||
|
|
||||||
m_a = self.obs_var * self.obs_count
|
|
||||||
m_b = batch_var * batch_count
|
|
||||||
M2 = m_a + m_b + delta**2 * self.obs_count * batch_count / total_count
|
|
||||||
new_var = M2 / total_count
|
|
||||||
|
|
||||||
self.obs_mean = new_mean
|
|
||||||
self.obs_var = new_var
|
|
||||||
self.obs_count = total_count
|
|
||||||
|
|
||||||
def _rollout(self, env_state, next_obs, next_done) -> tuple[Any, ...]:
|
def _rollout(self, env_state, next_obs, next_done) -> tuple[Any, ...]:
|
||||||
return self._rollout_jit(
|
return self._rollout_jit(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
|
|
@ -840,7 +722,7 @@ class PPOTrainer:
|
||||||
def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple:
|
def _step(self, env_state, next_obs, next_done, iteration: int) -> tuple:
|
||||||
if iteration == 1:
|
if iteration == 1:
|
||||||
self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}")
|
self.logger.log_non_interactive(f"Starting first rollout (JIT): {time.ctime()}")
|
||||||
self.logger.info(f"[_step] next_obs (in): {next_obs.shape}")
|
self.logger.debug(f"[_step] next_obs (in): {next_obs.shape}")
|
||||||
(
|
(
|
||||||
self.agent_state,
|
self.agent_state,
|
||||||
self.episode_stats,
|
self.episode_stats,
|
||||||
|
|
@ -850,12 +732,12 @@ class PPOTrainer:
|
||||||
self.key,
|
self.key,
|
||||||
next_env_state,
|
next_env_state,
|
||||||
) = self._rollout(env_state, next_obs, next_done)
|
) = self._rollout(env_state, next_obs, next_done)
|
||||||
self.logger.info(f"[_step] next_obs (post-rollout): {next_obs.shape}")
|
self.logger.debug(f"[_step] next_obs (post-rollout): {next_obs.shape}")
|
||||||
if iteration == 1:
|
if iteration == 1:
|
||||||
self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}")
|
self.logger.log_non_interactive(f"First rollout completed: {time.ctime()}")
|
||||||
|
|
||||||
storage = self._compute_gae(storage, next_obs, next_done)
|
storage = self._compute_gae(storage, next_obs, next_done)
|
||||||
self.logger.info(f"[_step] storage.obs (post-gae): {storage.obs.shape}")
|
self.logger.debug(f"[_step] storage.obs (post-gae): {storage.obs.shape}")
|
||||||
if iteration == 1:
|
if iteration == 1:
|
||||||
self.logger.log_non_interactive(f"Starting first PPO update (JIT): {time.ctime()}")
|
self.logger.log_non_interactive(f"Starting first PPO update (JIT): {time.ctime()}")
|
||||||
|
|
||||||
|
|
@ -931,15 +813,10 @@ class PPOTrainer:
|
||||||
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
|
self.logger.log_non_interactive(f"Initial reset started: {time.ctime()}")
|
||||||
|
|
||||||
env_state = self.env.reset(seed=self.experiment.seed)
|
env_state = self.env.reset(seed=self.experiment.seed)
|
||||||
|
|
||||||
next_obs = _convert_obs_dict_to_array_morphology(
|
next_obs = self.obs_processor(env_state.observations)
|
||||||
env_state.observations,
|
self.logger.debug(f"[train] next_obs: {next_obs.shape}")
|
||||||
self.morph_mode,
|
|
||||||
self.num_segments,
|
|
||||||
self.num_arms,
|
|
||||||
)
|
|
||||||
self.logger.info(f"[train] next_obs: {next_obs.shape}")
|
|
||||||
# TODO: next_obs = self.obs_processor(env_state.observations)
|
|
||||||
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
next_done = jnp.zeros(self.ppo.num_envs, dtype=jnp.bool_)
|
||||||
|
|
||||||
self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}")
|
self.logger.log_non_interactive(f"Initial reset completed: {time.ctime()}")
|
||||||
|
|
@ -954,9 +831,6 @@ class PPOTrainer:
|
||||||
env_state, next_obs, next_done, training_measurements, storage = self._step(
|
env_state, next_obs, next_done, training_measurements, storage = self._step(
|
||||||
env_state, next_obs, next_done, iteration=iteration
|
env_state, next_obs, next_done, iteration=iteration
|
||||||
)
|
)
|
||||||
self.logger.debug(f"[train] next_obs (post-step): {next_obs.shape}")
|
|
||||||
self._update_obs_stats(next_obs)
|
|
||||||
next_obs = _normalize_obs(next_obs, self.obs_mean, self.obs_var)
|
|
||||||
|
|
||||||
global_step += self.ppo.num_steps * self.ppo.num_envs
|
global_step += self.ppo.num_steps * self.ppo.num_envs
|
||||||
self._log(
|
self._log(
|
||||||
|
|
|
||||||
Reference in a new issue