def rollout(env, policy, num_episodes, rng):
"""Collect complete episodes from `policy(obs, rng) -> action` as a
Batch; `term` records `terminated`, never `truncated` (:numref:`sec_mdp`).
All sampling runs through the one numpy generator `rng`."""
cols, ep_ends = [[] for _ in range(5)], []
for _ in range(num_episodes):
obs, done = env.reset()[0], False
while not done:
act = policy(obs, rng)
next_obs, reward, terminated, truncated, _ = env.step(act)
done = terminated or truncated
for col, val in zip(cols, (obs, act, reward, next_obs,
float(terminated))):
col.append(val)
obs = next_obs
ep_ends.append(len(cols[0]))
obs, act, rew, next_obs, term = (np.asarray(c) for c in cols)
return Batch(obs, act, rew.astype(np.float32), next_obs,
term.astype(np.float32), np.asarray(ep_ends))