Created
March 31, 2021 14:55
-
-
Save ZJUGuoShuai/e43c463dc8163637350c31fd163b66b8 to your computer and use it in GitHub Desktop.
最基础的 Policy Gradient(PyTorch 实现)
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| import torch | |
| import torch.nn as nn | |
| from torch.distributions.categorical import Categorical | |
| from torch.optim import Adam | |
| import numpy as np | |
| import gym | |
| from gym.spaces import Discrete, Box | |
| def mlp(sizes, activation=nn.Tanh, output_activation=nn.Identity): | |
| # 构建一个前向神经网络 | |
| layers = [] | |
| for j in range(len(sizes)-1): | |
| act = activation if j < len(sizes)-2 else output_activation | |
| layers += [nn.Linear(sizes[j], sizes[j+1]), act()] | |
| return nn.Sequential(*layers) | |
| def train(env_name='CartPole-v0', hidden_sizes=[32], lr=1e-2, | |
| epochs=50, batch_size=5000, render=False): | |
| # 创建环境,检查状态/动作空间是否符合要求;得到 observation 和 action 的维度 | |
| env = gym.make(env_name) | |
| assert isinstance(env.observation_space, Box), \ | |
| "This example only works for envs with continuous state spaces." | |
| assert isinstance(env.action_space, Discrete), \ | |
| "This example only works for envs with discrete action spaces." | |
| obs_dim = env.observation_space.shape[0] | |
| n_acts = env.action_space.n | |
| # 创建策略的参数网络 | |
| logits_net = mlp(sizes=[obs_dim]+hidden_sizes+[n_acts]) | |
| # 从策略参数网络计算动作分布(策略) | |
| def get_policy(obs): | |
| logits = logits_net(obs) | |
| return Categorical(logits=logits) | |
| # 从策略采样动作 | |
| def get_action(obs): | |
| return get_policy(obs).sample().item() | |
| # 这个 loss 的梯度等于策略梯度的相反,因此用这个 loss SGD 相当于用策略梯度 SGA | |
| def compute_loss(obs, act, weights): | |
| logp = get_policy(obs).log_prob(act) | |
| return -(logp * weights).mean() | |
| # 创建优化器 | |
| optimizer = Adam(logits_net.parameters(), lr=lr) | |
| # 主训练函数(进行一次梯度更新) | |
| def train_one_epoch(): | |
| batch_obs = [] # for observations | |
| batch_acts = [] # for actions | |
| batch_weights = [] # for R(tau) weighting in policy gradient | |
| batch_rets = [] # for measuring episode returns | |
| batch_lens = [] # for measuring episode lengths | |
| # reset episode-specific variables | |
| obs = env.reset() # first obs comes from starting distribution | |
| done = False # signal from environment that episode is over | |
| ep_rews = [] # list for rewards accrued throughout ep | |
| # 用当前策略与环境交互,得到采样的轨迹 | |
| while True: | |
| # 环境渲染 | |
| if (not finished_rendering_this_epoch) and render: | |
| env.render() | |
| # 将初始 observation 保存 | |
| batch_obs.append(obs.copy()) | |
| # 在环境中依据策略做出动作,环境发生状态转移 | |
| act = get_action(torch.as_tensor(obs, dtype=torch.float32)) | |
| obs, rew, done, _ = env.step(act) | |
| # 保存动作和 reward | |
| batch_acts.append(act) | |
| ep_rews.append(rew) | |
| if done: | |
| # 如果 episode 结束,则记录整个 episode 相关信息 | |
| ep_ret, ep_len = sum(ep_rews), len(ep_rews) | |
| batch_rets.append(ep_ret) | |
| batch_lens.append(ep_len) | |
| # 每个 logprob(a|s) 对应的 weight 都是 R(tau) | |
| batch_weights += [ep_ret] * ep_len | |
| # 重置变量 | |
| obs, done, ep_rews = env.reset(), False, [] | |
| # 收集的样本足够一个 batch 了,则可以退出 | |
| if len(batch_obs) > batch_size: | |
| break | |
| # 策略网络的参数进行一次梯度更新 | |
| optimizer.zero_grad() | |
| batch_loss = compute_loss(obs=torch.as_tensor(batch_obs, dtype=torch.float32), | |
| act=torch.as_tensor(batch_acts, dtype=torch.int32), | |
| weights=torch.as_tensor(batch_weights, | |
| dtype=torch.float32 | |
| ) | |
| ) | |
| batch_loss.backward() | |
| optimizer.step() | |
| return batch_loss, batch_rets, batch_lens | |
| # 训练循环 | |
| for i in range(epochs): | |
| batch_loss, batch_rets, batch_lens = train_one_epoch() | |
| print('epoch: %3d \t loss: %.3f \t return: %.3f \t ep_len: %.3f'% | |
| (i, batch_loss, np.mean(batch_rets), np.mean(batch_lens))) | |
| if __name__ == '__main__': | |
| import argparse | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument('--env_name', '--env', type=str, default='CartPole-v0') | |
| parser.add_argument('--render', action='store_true') | |
| parser.add_argument('--lr', type=float, default=1e-2) | |
| args = parser.parse_args() | |
| print('\nUsing simplest formulation of policy gradient.\n') | |
| train(env_name=args.env_name, render=args.render, lr=args.lr) |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment