def train_dqn(seed, qnet, use_target=True, step=None):
"""DQN on CartPole; yields (env step, episode return, max_a Q(s0, a))."""
step = q_step if step is None else step
rng, env = np.random.default_rng(seed), gym.make('CartPole-v1')
target = make_qnet()
target.load_state_dict(qnet.state_dict())
opt = torch.optim.Adam(qnet.parameters(), lr=lr)
buffer, s0 = ReplayBuffer(buffer_size, 4), np.zeros(4, np.float32)
obs, ep_return = env.reset(seed=seed)[0], 0.0
sync = sync_every if use_target else 1
for t in range(1, num_env_steps + 1):
a = d2l.epsilon_greedy(q_values(qnet, obs), epsilon(t), rng)
next_obs, rew, terminated, truncated, _ = env.step(a)
buffer.add(obs, a, rew, next_obs, float(terminated))
obs, ep_return = next_obs, ep_return + rew
if terminated or truncated:
yield t, ep_return, q_values(qnet, s0).max()
obs, ep_return = env.reset()[0], 0.0
if len(buffer) >= warmup and t % train_freq == 0:
step(qnet, target, opt, buffer.sample(batch_size, rng))
if t % sync == 0:
target.load_state_dict(qnet.state_dict())