强化学习实战:策略梯度中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,往往是让训练重回正轨的有效手段。

Logo

更多推荐