def train_ppo(seed, ac, use_clip=True, trace=None):
"""Freeze the advantages and the collecting policy's log-probs, then
spend num_epochs surrogate passes; GAE(0.95) is the default."""
rng, env = np.random.default_rng(seed), gym.make('CartPole-v1')
env.reset(seed=seed)
for _ in range(num_updates):
batch = d2l.rollout(env, ac.act, batch_episodes, rng)
for _ in range(critic_steps): # fresh lambda-return target, per pass
fit_value(ac, batch.obs, batch.gae(ac.value_np, gamma, lam)
+ ac.value_np(batch.obs))
adv = d2l.normalize(batch.gae(ac.value_np, gamma, lam))
logp_old = ac.log_prob_np(batch.obs, batch.act)
d = ppo_epochs(ac, batch, adv, logp_old, epsilon_clip, num_epochs,
entropy_coef, use_clip)
if trace is not None:
trace.append(d)
yield (float(batch.episode_returns().mean()), *d.mean(0), *d[-1])