MCPcopy Create free account
hub / github.com/BIT-MCS/DRL-eFresh / train

Function train

utils/main_mp_github.py:57–103  ·  view source on GitHub ↗
(rank, agent, config)

Source from the content-addressed store, hash-verified

55
56
57def train(rank, agent, config):
58 env = gym.make("Seaquest-v0")
59 torch.manual_seed(config.seed+rank)
60 env.seed(config.seed+rank)
61
62 policy = Policy(agent=agent)
63 optimizer = optim.Adam(policy.parameters(), lr=1e-3)
64 running_reward = 10.0
65
66 # NOTE: I am using a different update mechanism as of now (REINFORCE vs. A3C).
67 for i_episode in range(config.num_episodes):
68 observation = env.reset()
69 # resets hidden states, otherwise the comp. graph history spans episodes
70 # and relies on freed buffers.
71 agent.reset() # NOTE: This may be problematic across processes.
72 ep_reward = 0
73
74 # Stash model in case of crash.
75 if i_episode % config.save_model_interval == 0 and i_episode > 0:
76 torch.save(agent.state_dict(), f"./models/agent-{i_episode}-{rank}.pt")
77
78 for t in range(config.max_steps):
79 action = policy(observation)
80 reward = 0.0
81 for _ in range(config.num_repeat_action):
82 if config.render:
83 env.render()
84 observation, _reward, done, _ = env.step(action)
85 reward += _reward
86 if done:
87 break
88 policy.rewards.append(reward)
89 ep_reward += reward
90 if done:
91 running_reward = 0.05 * ep_reward + (1 - 0.05) * running_reward
92 finish_episode(optimizer, policy, config)
93 if i_episode % config.log_interval == 0:
94 print(
95 f"Episode {i_episode}-{rank}\tLast reward: {ep_reward:.2f}\tAverage reward: {running_reward:.2f}"
96 )
97 if running_reward > config.reward_threshold:
98 print(
99 f"Solved! Running reward is now {running_reward} and "
100 f"the last episode runs to {t} time steps!"
101 )
102 break
103 env.close()
104
105if __name__ == "__main__":
106 parser = argparse.ArgumentParser()

Callers

nothing calls this directly

Calls 5

PolicyClass · 0.85
finish_episodeFunction · 0.85
renderMethod · 0.80
resetMethod · 0.45
stepMethod · 0.45

Tested by

no test coverage detected