强化学习必备技巧:Policy Gradient中Categorical分布的5个关键用法
强化学习实战:策略梯度中Categorical分布的五个核心应用与避坑指南
在强化学习的探索与实践中,策略梯度方法因其直接优化策略参数的能力而备受青睐。无论是经典的REINFORCE算法,还是后来的Actor-Critic框架,其核心都离不开一个关键组件:策略的概率分布表示。对于离散动作空间,这个分布通常就是分类分布。许多开发者虽然能调用torch.distributions.Categorical完成采样,但在实际构建稳定、高效的智能体时,往往会遇到梯度消失、训练不稳定或采样偏差等问题。这些问题,往往源于对Categorical分布的理解停留在表面,未能深入其与策略梯度理论结合的细节。
本文将从一个实践者的视角出发,抛开泛泛而谈的API介绍,聚焦于策略梯度算法中Categorical分布的五个关键且具体的应用场景。我们会结合PyTorch的torch.distributions.Categorical,不仅展示“如何用”,更深入探讨“为何这样用”,以及在实际编码和调参中可能遇到的“坑”和“避坑”策略。无论你是在构建一个玩Atari游戏的智能体,还是设计一个解决组合优化问题的策略网络,理解这些用法都将帮助你写出更健壮、更高效的代码。
1. 策略表示:从Logits到可微分的概率分布
在策略网络中,我们通常不会直接输出概率值,而是输出未经归一化的Logits。这背后有数值稳定性和优化便利性的双重考量。直接使用Categorical(probs=...)看似直观,但在反向传播时可能面临梯度问题,尤其是在概率值接近0或1的边界区域。
1.1 为何首选Logits参数?
使用logits参数初始化Categorical分布,是强化学习中的标准做法。PyTorch内部会使用log_softmax进行归一化,这个过程在数值上更加稳定。
import torch
import torch.nn as nn
import torch.nn.functional as F
class PolicyNetwork(nn.Module):
def __init__(self, obs_dim, act_dim):
super().__init__()
self.fc1 = nn.Linear(obs_dim, 64)
self.fc2 = nn.Linear(64, 64)
self.head = nn.Linear(64, act_dim) # 输出logits
def forward(self, obs):
x = torch.tanh(self.fc1(obs))
x = torch.tanh(self.fc2(x))
logits = self.head(x) # 未归一化的分数
return logits
# 假设环境有4个离散动作
policy_net = PolicyNetwork(obs_dim=8, act_dim=4)
state = torch.randn(1, 8) # 一个批次的观测
logits = policy_net(state)
# 正确做法:使用logits参数
action_dist = torch.distributions.Categorical(logits=logits)
# 避免做法:手动softmax后再传入probs
# probs = F.softmax(logits, dim=-1) # 可能导致梯度问题
# action_dist = torch.distributions.Categorical(probs=probs)
注意:
Categorical类在接收logits参数时,内部会自动处理数值稳定性,例如在计算log_prob时使用log_softmax,这比先计算softmax再取对数要稳定得多,尤其是在logits值较大或较小时。
1.2 处理批量数据与多维动作空间
在实际训练中,我们经常处理的是批量数据。Categorical完美支持批量操作,但需要确保维度正确。对于更复杂的动作空间(例如多个独立的离散动作),我们需要创建多个独立的Categorical分布。
# 批量处理示例:批量大小为32,动作维度为5
batch_logits = torch.randn(32, 5)
batch_dist = torch.distributions.Categorical(logits=batch_logits)
# 此时batch_dist.batch_shape为torch.Size([32]),event_shape为torch.Size([])
# 采样:为批次中的每个状态采样一个动作
batch_actions = batch_dist.sample() # shape: [32]
# 计算该批次中每个动作对应的对数概率
batch_log_probs = batch_dist.log_prob(batch_actions) # shape: [32]
# 多维离散动作示例(例如,需要同时选择“移动方向”和“攻击类型”)
logits_direction = torch.randn(32, 4) # 4个方向
logits_attack = torch.randn(32, 3) # 3种攻击
dist_direction = torch.distributions.Categorical(logits=logits_direction)
dist_attack = torch.distributions.Categorical(logits=logits_attack)
action_dir = dist_direction.sample()
action_atk = dist_attack.sample()
# 总的对数概率是各个独立动作分布对数概率之和(假设动作独立)
log_prob = dist_direction.log_prob(action_dir) + dist_attack.log_prob(action_atk)
2. 动作采样:探索与利用的权衡艺术
采样是策略执行的核心。sample()方法看似简单,但其背后的随机性直接决定了智能体的探索行为。如何控制探索的程度,是策略梯度算法成败的关键之一。
2.1 基础采样与确定性测试
在训练阶段,我们依赖采样来进行探索;而在评估或部署阶段,我们往往选择概率最高的动作(贪婪策略)以获得确定性行为。
def select_action(state, training=True):
logits = policy_net(state)
dist = torch.distributions.Categorical(logits=logits)
if training:
# 训练时:采样,引入探索
action = dist.sample()
else:
# 评估时:选择概率最高的动作,追求稳定表现
action = torch.argmax(logits, dim=-1) # 等价于 dist.probs.argmax()
return action.item() if state.dim() == 1 else action
2.2 利用熵正则化控制探索强度
单纯的采样有时会导致探索不足,智能体容易陷入局部最优。一个常见的技巧是在损失函数中加入策略熵的负值作为正则项,鼓励分布更均匀(即探索更多)。
def compute_loss(trajectory_batch):
states, actions, returns = trajectory_batch
logits = policy_net(states)
dist = torch.distributions.Categorical(logits=logits)
# 核心策略梯度损失:负对数概率乘以优势函数
log_probs = dist.log_prob(actions)
policy_loss = -(log_probs * returns).mean()
# 熵正则化:增加熵意味着鼓励探索
# beta是控制探索强度的超参数,通常设为0.01左右
entropy_bonus = dist.entropy().mean()
beta = 0.01
total_loss = policy_loss - beta * entropy_bonus
return total_loss, entropy_bonus.item()
我们可以通过一个简单的表格来理解不同熵值下分布的特点:
| 熵值范围 | 分布形态 | 探索倾向 | 典型场景 |
|---|---|---|---|
| 接近0 | 接近one-hot(确定性) | 低,偏向利用 | 训练后期,策略收敛时 |
| 中等 | 部分动作概率较高 | 中等 | 训练中期 |
| 较高(接近最大熵) | 接近均匀分布 | 高,积极探索 | 训练初期,或需要大量探索的环境 |
2.3 采样中的常见陷阱:torch.multinomial vs Categorical.sample()
有时开发者会尝试手动实现采样,例如使用torch.multinomial。虽然功能相似,但直接使用Categorical.sample()是更优选择,因为它与分布对象紧密绑定,能确保后续log_prob计算的正确性。
logits = torch.tensor([1.0, 2.0, 1.0])
dist = torch.distributions.Categorical(logits=logits)
# 推荐:使用分布内置方法
action = dist.sample()
correct_log_prob = dist.log_prob(action) # 正确
# 不推荐:手动采样,容易导致不一致
probs = F.softmax(logits, dim=-1)
manual_action = torch.multinomial(probs, num_samples=1).squeeze()
# 问题:如果后续用dist.log_prob(manual_action),虽然数学上可能对,但破坏了分布的抽象封装
3. 对数概率计算:策略梯度的基石
策略梯度定理告诉我们,参数的更新方向是期望回报关于策略对数概率的梯度。因此,准确、高效地计算log_prob是策略梯度算法的核心。
3.1 理解log_prob的输入与输出
dist.log_prob(action)返回的是在给定分布下,采样到特定动作action的对数概率。这里的关键是,action必须是代表类别索引的整数张量。
# 单个动作
logits_single = torch.tensor([0.5, 1.5, 0.0])
dist_single = torch.distributions.Categorical(logits=logits_single)
action = torch.tensor(1)
log_p = dist_single.log_prob(action) # 计算动作1的对数概率
print(f"Log prob of action 1: {log_p.item():.4f}")
# 批量计算:高效向量化操作
logits_batch = torch.tensor([[0.5, 1.5, 0.0], [1.0, 0.5, 0.5]])
actions_batch = torch.tensor([1, 2])
dist_batch = torch.distributions.Categorical(logits=logits_batch)
log_probs_batch = dist_batch.log_prob(actions_batch)
print(f"Batch log probs: {log_probs_batch}") # 分别对应两个状态-动作对
提示:在收集轨迹时,务必在采样动作的同一时刻保存对应的对数概率。不要先采样,然后在后续步骤中根据新的网络参数重新计算旧状态-动作对的
log_prob,因为网络参数已经更新,分布发生了变化,计算出的梯度将是错误的。
3.2 策略梯度损失函数的实现细节
在REINFORCE或Actor-Critic算法中,策略损失通常表示为优势函数与对数概率的乘积的负期望。使用PyTorch实现时,需要注意detach()的使用。
def reinforce_update(trajectory):
"""
trajectory: 包含状态、动作、回报的列表
"""
states = torch.stack([s for s, _, _ in trajectory])
actions = torch.tensor([a for _, a, _ in trajectory])
returns = torch.tensor([r for _, _, r in trajectory])
# 通常会对returns进行归一化,以稳定训练
returns = (returns - returns.mean()) / (returns.std() + 1e-8)
# 前向传播,获取当前策略下的分布
logits = policy_net(states)
dist = torch.distributions.Categorical(logits=logits)
log_probs = dist.log_prob(actions)
# 策略损失: -Σ (log_prob * G_t)
# 关键:returns在这里应被视为常数,不参与策略参数求导,因此需要detach
policy_loss = -(log_probs * returns.detach()).mean()
optimizer.zero_grad()
policy_loss.backward()
optimizer.step()
这里returns.detach()至关重要。它告诉PyTorch,在计算policy_loss对策略网络参数的梯度时,将returns视为一个固定的标量权重,而不是一个需要梯度的变量。梯度只会通过log_probs流回网络,这正是策略梯度定理所要求的。
3.3 处理带掩码(Mask)的动作空间
在某些环境中(如某些游戏或任务规划),并非所有动作在所有状态下都有效。我们需要将无效动作的概率置零,并重新归一化。Categorical不直接支持掩码,但我们可以通过操作logits来实现。
def create_masked_distribution(logits, action_mask):
"""
logits: [batch_size, num_actions]
action_mask: [batch_size, num_actions], 有效动作为1,无效为0
"""
# 将无效动作的logits设置为一个极小的负数,这样softmax后概率接近0
VERY_NEGATIVE = -1e10
masked_logits = logits + torch.log(action_mask.float() + 1e-8)
# 更鲁棒的做法是直接赋值一个非常大的负数
# masked_logits = logits.masked_fill(action_mask == 0, VERY_NEGATIVE)
dist = torch.distributions.Categorical(logits=masked_logits)
return dist
# 示例:状态1有3个有效动作,状态2有2个有效动作
batch_logits = torch.randn(2, 5)
mask = torch.tensor([[1, 1, 1, 0, 0],
[0, 0, 1, 1, 0]]) # 0表示无效
masked_dist = create_masked_distribution(batch_logits, mask)
sampled_actions = masked_dist.sample()
# 采样到的动作将只会在有效动作中产生
4. 梯度流与数值稳定性:避免训练崩溃的实用技巧
策略梯度方法以训练不稳定著称。很多时候,不稳定并非源于算法本身,而是由于实现细节上的疏忽,特别是在处理Categorical分布时。
4.1 检查梯度:log_prob的梯度爆炸与消失
log_prob的梯度大小直接影响到参数更新的幅度。如果概率非常接近0,其对数会趋向负无穷,梯度可能变得异常大(爆炸)或出现NaN。
# 一个可能导致问题的场景
problematic_logits = torch.tensor([1000.0, 0.0, 0.0]) # 一个logits值过大
dist = torch.distributions.Categorical(logits=problematic_logits)
# 此时softmax(1000,0,0)会导致第一个概率接近1,其余接近0
# 计算动作0的log_prob是没问题的,但计算动作1的log_prob会得到一个非常大的负数
prob = dist.probs
print(f"Probabilities: {prob}")
print(f"Log prob of action 1: {dist.log_prob(torch.tensor(1))}") # 可能是一个很大的负数
# 解决方案:对logits进行适当的缩放或归一化(例如,减去最大值)
stable_logits = problematic_logits - problematic_logits.max()
dist_stable = torch.distributions.Categorical(logits=stable_logits)
print(f"Stable probs: {dist_stable.probs}")
print(f"Stable log prob of action 1: {dist_stable.log_prob(torch.tensor(1))}")
在实践中,保持logits在一个合理的数值范围内(例如,通过网络最后一层不使用过大的权重初始化,或添加批归一化)是避免此类问题的关键。
4.2 使用torch.where处理边缘情况
在计算损失时,有时会遇到极端情况。例如,当优势函数估计值(或回报)非常大时,与log_prob相乘可能导致巨大的损失值。一个实用的技巧是使用torch.where进行裁剪或条件处理。
def clipped_policy_loss(log_probs, advantages, clip_ratio=0.2):
"""
近似PPO风格的裁剪损失,用于稳定训练。
"""
ratio = torch.exp(log_probs - log_probs.detach()) # 简单示例,实际PPO需要旧策略概率
clipped_advantages = torch.where(advantages > 0,
(1 + clip_ratio) * advantages,
(1 - clip_ratio) * advantages)
# 损失取最小值,实现裁剪效果
loss = -torch.min(ratio * advantages, clipped_advantages).mean()
return loss
4.3 监控训练:熵与KL散度
持续监控策略分布的变化是调试训练过程的重要手段。除了熵,计算当前策略与旧策略或某个参考策略之间的KL散度也很有用。
def compute_kl_divergence(new_logits, old_logits):
"""
计算两个Categorical分布之间的KL散度 D_KL(old || new)
"""
old_dist = torch.distributions.Categorical(logits=old_logits)
new_dist = torch.distributions.Categorical(logits=new_logits)
# KL散度 = Σ old_prob * (log(old_prob) - log(new_prob))
# 也可以直接用PyTorch的kl_divergence,但注意参数顺序
kl = torch.distributions.kl.kl_divergence(old_dist, new_dist)
return kl.mean()
# 在训练循环中监控
current_entropy = dist.entropy().mean().item()
kl = compute_kl_divergence(current_logits, old_logits.detach())
if kl > 0.01: # 如果策略更新步幅过大
print(f"Warning: KL divergence ({kl:.4f}) is large, consider reducing learning rate.")
5. 超越基础:高级模式与性能优化
当算法和网络变得复杂时,对Categorical分布的高效使用提出了更高要求。以下是一些进阶技巧。
5.1 与torch.distributions.Independent结合处理多维动作
如前所述,对于多维离散动作,我们需要多个Categorical分布。torch.distributions.Independent可以帮助我们将多个独立分布的乘积视为一个联合分布,从而简化log_prob的计算。
from torch.distributions import Independent
# 假设一个动作由两个独立部分构成:移动(3种)和攻击(2种)
logits_move = torch.randn(32, 3)
logits_attack = torch.randn(32, 2)
dist_move = torch.distributions.Categorical(logits=logits_move)
dist_attack = torch.distributions.Categorical(logits=logits_attack)
# 传统方法:分别采样,对数概率相加
action_move = dist_move.sample()
action_attack = dist_attack.sample()
log_prob_manual = dist_move.log_prob(action_move) + dist_attack.log_prob(action_attack)
# 使用Independent:创建联合分布
# 首先将两个分布堆叠成一个“多变量”分布(事件形状为[2])
joint_distribution = Independent(torch.stack([dist_move, dist_attack], dim=-1), 1)
# 采样得到形状为[32, 2]的动作
joint_action = joint_distribution.sample() # 第一列是move,第二列是attack
# 直接计算联合对数概率
joint_log_prob = joint_distribution.log_prob(joint_action)
print(torch.allclose(log_prob_manual, joint_log_prob)) # 应为True
5.2 自定义分布以满足特殊需求
有时标准Categorical无法满足需求,例如需要实现带温度参数(Temperature)的采样,常用于软策略或知识蒸馏。我们可以通过继承torch.distributions.Categorical来轻松实现。
class CategoricalWithTemperature(torch.distributions.Categorical):
def __init__(self, logits=None, probs=None, temperature=1.0):
super().__init__(logits=logits, probs=probs)
self.temperature = temperature
def sample(self, sample_shape=torch.Size()):
# 应用温度缩放:logits / temperature
if self.temperature != 1.0:
scaled_logits = self.logits / self.temperature
# 需要创建一个临时的分布进行采样
temp_dist = torch.distributions.Categorical(logits=scaled_logits)
return temp_dist.sample(sample_shape)
else:
return super().sample(sample_shape)
def log_prob(self, value):
# 注意:对数概率也需要用相同的温度缩放进行校正
if self.temperature != 1.0:
scaled_logits = self.logits / self.temperature
temp_dist = torch.distributions.Categorical(logits=scaled_logits)
return temp_dist.log_prob(value)
else:
return super().log_prob(value)
# 使用示例
logits = torch.tensor([1.0, 2.0, 3.0])
high_temp_dist = CategoricalWithTemperature(logits=logits, temperature=2.0)
print(f"High temp probs: {high_temp_dist.probs}") # 分布更均匀
low_temp_dist = CategoricalWithTemperature(logits=logits, temperature=0.5)
print(f"Low temp probs: {low_temp_dist.probs}") # 分布更尖锐
5.3 在向量化环境中的高效应用
现代强化学习库(如gym.vector或自定义环境)支持多个环境并行运行。Categorical的批量处理能力在这里大放异彩。
# 假设我们在8个并行环境中运行
num_envs = 8
obs_batch = envs.reset() # 形状假设为 [8, obs_dim]
logits_batch = policy_net(obs_batch) # [8, num_actions]
dist_batch = torch.distributions.Categorical(logits=logits_batch)
actions_batch = dist_batch.sample() # [8]
# 执行动作,获取下一个状态、奖励等
next_obs_batch, rewards_batch, dones_batch, _ = envs.step(actions_batch.cpu().numpy())
# 计算用于更新的对数概率(只针对未结束的环境)
if not all(dones_batch):
# 可能只需要为未结束的状态计算log_prob
mask = ~torch.tensor(dones_batch)
active_log_probs = dist_batch.log_prob(actions_batch)[mask]
掌握这些从基础到高级的用法,意味着你不仅能在强化学习项目中正确使用Categorical分布,更能理解其背后的原理,并能够根据具体任务灵活调整和优化。最终,写出既稳定又高效的代码,让你的智能体学习得更快、更好。在实际项目中,我习惯在训练循环开始时记录初始策略的熵,并在每次更新后监控其变化,这比单纯看回报曲线更能揭示策略的探索状态。当熵值下降过快时,适当调高熵正则项的系数beta,往往是让训练重回正轨的有效手段。
更多推荐

所有评论(0)