模仿学习2.4:生成式对抗模仿学习GAIL
概念
是一种让计算机通过观察专家示范来学会完成任务的机器学习方法。
可以把它想象成一个"表演家"和"评论家"之间的博弈:
-
表演家(生成器):一个新手,试图模仿专家的动作。
-
评论家(判别器):一个考官,火眼金睛地分辨哪些动作是专家的,哪些是新手模仿的。
整个过程是一个"最小-最大"博弈。智能体(生成器)的目标是最小化自己与专家的差距,而判别器的目标是最大化自己分辨真伪的能力。随着博弈的进行,新手为了骗过考官,会模仿得越来越像,最终成为一个能像专家一样熟练执行任务的"行家"。
生成器(generator)和判别器(discriminator)各是一个神经网络。生成器负责生成假的样本,而判别器负责判定一个样本是真是假。判别器像一个不断进化的"专家鉴定师",它被训练去区分哪些"状态-动作对"来自真正的专家,哪些来自正在学习的智能体。智能体(生成器)则在环境中行动。它的"奖励"不是来自环境,而是来自判别器的打分。如果判别器认为它的行为像专家,它就得到高分。为了持续获得高分,智能体必须不断调整自己的策略,让自己更像专家。
工作流程
-
准备专家数据
收集专家完成任务时的状态-动作对(或完整轨迹),作为模仿的目标。 -
初始化两个网络
-
生成器网络:根据状态输出动作。
-
判别器网络:输入状态-动作对,输出该对来自专家的概率。
-
-
迭代训练
在每个训练轮次中:-
生成策略数据:用当前策略与环境交互,采集一批状态-动作对。
-
采样专家数据:从专家数据集中随机抽取相同数量的状态-动作对。
-
训练判别器
将策略数据和专家数据混合,训练判别器区分二者(专家数据标为1,策略数据标为0)。损失函数通常是二分类交叉熵。 -
计算伪奖励
用判别器的输出为策略生成奖励信号(例如-log(1 - D(s,a))或log D(s,a)),使策略能通过强化学习算法优化。 -
更新策略
使用任意强化学习算法(如PPO、TRPO)根据伪奖励更新策略,目标是让判别器更难区分策略数据与专家数据。
-
-
重复
交替训练判别器和策略,直至策略的行为与专家足够接近。
代码实现
1.PPO训练专家数据(见模仿学习2.1)
2.GAIL_cartpole.py
import paddle
import paddle.nn as nn
import paddle.nn.functional as F
import gymnasium as gym
import numpy as np
import matplotlib.pyplot as plt
from tqdm import tqdm
from paddle.distribution import Categorical
# -------------------- 网络定义 --------------------
# 策略网络(Actor)
class PolicyNet(nn.Layer):
def __init__(self, state_dim, hidden_dim, action_dim):
super(PolicyNet, self).__init__()
self.fc1 = nn.Linear(state_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, action_dim)
def forward(self, x):
x = F.relu(self.fc1(x))
return self.fc2(x) # logits
# 判别器网络(GAIL核心)
class Discriminator(nn.Layer):
def __init__(self, state_dim, hidden_dim, action_dim):
super(Discriminator, self).__init__()
self.action_dim = action_dim
self.fc1 = nn.Linear(state_dim + action_dim, hidden_dim)
self.fc2 = nn.Linear(hidden_dim, hidden_dim)
self.fc3 = nn.Linear(hidden_dim, 1)
def forward(self, state, action):
# 动作one-hot编码
action_one_hot = F.one_hot(action, num_classes=self.action_dim)
action_one_hot = paddle.squeeze(action_one_hot, axis=1)
# 拼接状态+动作
x = paddle.concat([state, action_one_hot], axis=1)
x = F.relu(self.fc1(x))
x = F.relu(self.fc2(x))
# 输出概率:越接近1=专家,越接近0=智能体
return F.sigmoid(self.fc3(x))
# -------------------- 专家数据采样 --------------------
def sample_expert_data(n_episode, env, model):
states = []
actions = []
for episode in range(n_episode):
reset_result = env.reset()
state = reset_result[0] if isinstance(reset_result, tuple) else reset_result
done = False
while not done:
state_tensor = paddle.to_tensor(state, dtype='float32').unsqueeze(0)
logits = model(state_tensor)
probs = F.softmax(logits, axis=-1)
action = paddle.argmax(probs, axis=-1).item()
states.append(state)
actions.append(action)
step_result = env.step(action)
if len(step_result) == 5:
next_state, reward, terminated, truncated, _ = step_result
done = terminated or truncated
else:
next_state, reward, done, _ = step_result
state = next_state
return np.array(states), np.array(actions)
# -------------------- GAIL智能体 --------------------
class GAIL:
def __init__(self, state_dim, hidden_dim, action_dim, lr, gamma=0.98):
self.gamma = gamma
# 策略网络 + 优化器
self.policy = PolicyNet(state_dim, hidden_dim, action_dim)
self.policy_optim = paddle.optimizer.Adam(
parameters=self.policy.parameters(), learning_rate=lr
)
# 判别器网络 + 优化器
self.discriminator = Discriminator(state_dim, hidden_dim, action_dim)
self.disc_optim = paddle.optimizer.Adam(
parameters=self.discriminator.parameters(), learning_rate=lr
)
# 选择动作
def take_action(self, state, deterministic=False):
state = paddle.to_tensor([state], dtype="float32")
logits = self.policy(state)
probs = F.softmax(logits, axis=-1)
if deterministic:
return paddle.argmax(probs, axis=-1).item()
else:
return Categorical(probs).sample([1]).item()
# 判别器损失:区分专家/智能体
def disc_loss(self, s_expert, a_expert, s_agent, a_agent):
d_expert = self.discriminator(s_expert, a_expert)
d_agent = self.discriminator(s_agent, a_agent)
# 二分类交叉熵损失
loss_expert = -paddle.log(d_expert + 1e-8).mean()
loss_agent = -paddle.log(1 - d_agent + 1e-8).mean()
return loss_expert + loss_agent
# 策略损失:最大化判别器奖励(骗过判别器)
def policy_loss(self, s, a):
d = self.discriminator(s, a)
# 奖励 = -log(1-D)
reward = -paddle.log(1 - d + 1e-8)
logits = self.policy(s)
log_probs = F.log_softmax(logits, axis=-1)
action_log_probs = paddle.take_along_axis(log_probs, a, axis=1)
# 策略梯度:最大化奖励
loss = -(action_log_probs * reward.detach()).mean()
return loss
# 单步训练
def learn(self, s_expert, a_expert, s_agent, a_agent):
# 转tensor
s_expert = paddle.to_tensor(s_expert, dtype="float32")
a_expert = paddle.to_tensor(a_expert, dtype="int64").unsqueeze(1)
s_agent = paddle.to_tensor(s_agent, dtype="float32")
a_agent = paddle.to_tensor(a_agent, dtype="int64").unsqueeze(1)
# 更新判别器
d_loss = self.disc_loss(s_expert, a_expert, s_agent, a_agent)
self.disc_optim.clear_grad()
d_loss.backward()
self.disc_optim.step()
# 更新策略
p_loss = self.policy_loss(s_agent, a_agent)
self.policy_optim.clear_grad()
p_loss.backward()
self.policy_optim.step()
return d_loss.item(), p_loss.item()
# -------------------- 测试函数 --------------------
def test_agent(agent, env, n_episode, deterministic=True):
return_list = []
for _ in range(n_episode):
reset_result = env.reset()
state = reset_result[0] if isinstance(reset_result, tuple) else reset_result
episode_return = 0
done = False
while not done:
action = agent.take_action(state, deterministic=deterministic)
step_result = env.step(action)
if len(step_result) == 5:
next_state, reward, terminated, truncated, _ = step_result
done = terminated or truncated
else:
next_state, reward, done, _ = step_result
state = next_state
episode_return += reward
return_list.append(episode_return)
return np.mean(return_list)
def test_expert(env, model, n_episodes=10):
returns = []
for _ in range(n_episodes):
reset_result = env.reset()
state = reset_result[0] if isinstance(reset_result, tuple) else reset_result
episode_return = 0
done = False
while not done:
state_tensor = paddle.to_tensor(state, dtype='float32').unsqueeze(0)
logits = model(state_tensor)
probs = F.softmax(logits, axis=-1)
action = paddle.argmax(probs, axis=-1).item()
step_result = env.step(action)
if len(step_result) == 5:
state, reward, terminated, truncated, _ = step_result
done = terminated or truncated
else:
state, reward, done, _ = step_result
episode_return += reward
returns.append(episode_return)
return np.mean(returns)
# -------------------- 主程序 --------------------
if __name__ == "__main__":
env_name = 'CartPole-v1'
env = gym.make(env_name)
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.n
hidden_dim = 128
lr = 1e-4
# 加载专家模型
expert_actor = PolicyNet(state_dim, hidden_dim, action_dim)
try:
expert_actor.set_state_dict(paddle.load("net_ppo.pdparams"))
print("专家模型加载成功")
except Exception as e:
print("加载失败,请检查文件路径:", e)
exit()
# 测试专家
expert_score = test_expert(env, expert_actor)
print(f"专家平均回报: {expert_score:.2f}")
# 采样专家数据
expert_s, expert_a = sample_expert_data(100, env, expert_actor)
print(f"专家数据量: {len(expert_s)}")
# 初始化GAIL
gail_agent = GAIL(state_dim, hidden_dim, action_dim, lr)
# 训练参数
total_iterations = 20000
batch_size = 64
test_freq = 200
test_returns = []
# 训练循环
print("\n开始GAIL训练...")
with tqdm(total=total_iterations) as pbar:
for i in range(total_iterations):
# 1. 智能体收集轨迹
s_batch, a_batch = [], []
while len(s_batch) < batch_size:
reset_result = env.reset()
state = reset_result[0] if isinstance(reset_result, tuple) else reset_result
done = False
while not done and len(s_batch) < batch_size:
s_batch.append(state)
action = gail_agent.take_action(state)
a_batch.append(action)
step_result = env.step(action)
if len(step_result) == 5:
next_state, _, terminated, truncated, _ = step_result
done = terminated or truncated
else:
next_state, _, done, _ = step_result
state = next_state
# 2. 随机采样专家数据
idx = np.random.choice(len(expert_s), batch_size, replace=False)
s_exp = expert_s[idx]
a_exp = expert_a[idx]
# 3. 训练GAIL
d_loss, p_loss = gail_agent.learn(s_exp, a_exp, s_batch, a_batch)
# 4. 测试
if (i + 1) % test_freq == 0:
test_r = test_agent(gail_agent, env, 5)
test_returns.append(test_r)
pbar.set_postfix({
"D_loss": f"{d_loss:.3f}",
"P_loss": f"{p_loss:.3f}",
"Return": f"{test_r:.1f}"
})
pbar.update(1)
# 绘图
plt.plot([i*test_freq for i in range(len(test_returns))], test_returns, marker='o')
plt.xlabel("Training Iteration")
plt.ylabel("Average Return")
plt.title(f"GAIL on {env_name} (Expert: {expert_score:.1f})")
plt.grid(True)
plt.show()
env.close()

更多推荐
所有评论(0)