强化学习行为克隆中狄拉克δ分布拟合正态分布的问题排查
问题背景
我正尝试在连续动作空间的Actor-Critic智能体上应用行为克隆方法,目标是用正态分布拟合狄拉克δ分布p_bc(a|x)(其中a为动作,x为状态)。我用PyTorch构建了一个Actor网络,试图学习给定x对应的均值μ(x)和标准差σ(x),已测试MSE和NLL损失函数,想要尝试KL损失但未找到狄拉克δ分布相关文档,不确定PyTorch是否有对应分布模块。
代码实现
网络结构
class DistributionFitModel(nn.Module): def __init__(self): super().__init__() self.model = nn.Sequential( nn.Linear(1, 16), nn.Linear(16, 32), nn.Linear(32, 16) ) self.mu = nn.Sequential(nn.Linear(16, 1), nn.ReLU()) self.logvar = nn.Sequential(nn.Linear(16, 1), nn.ReLU()) def forward(self, x): x = self.model(x) mu = self.mu(x) logvar = self.logvar(x) std = torch.exp(0.5 * logvar) eps = torch.randn_like(std) return eps.mul(std).add_(mu) def mu_sigma(self, x): x = self.model(x) return self.mu(x), torch.exp(0.5 * self.logvar(x))
我采用logvar而非直接用var,参考了VAE的实现方式。
数据生成
keys = [(0, -1), (1, 1), (2, -1), (3, 0), (4, 1), (5, 0)] ys = [] xs = [] L = 1000 for x_mul, y_mul in keys: ys += (np.ones(1000) * y_mul).tolist() xs += (np.ones(1000) * x_mul).tolist()
为每个x生成1000个对应动作样本(如x=0对应动作a=-1)。
数据集与数据加载器
class DistDataset(Dataset): def __init__(self, x, y): super().__init__() self.x = x self.y = y def __len__(self): return len(self.x) def __getitem__(self, item): return torch.tensor(self.x[item]), torch.tensor(self.y[item]) dataset = DistDataset(xs, ys) dataloader = DataLoader(dataset, batch_size=512, shuffle=True)
训练循环
for _ in range(1000): running_loss = 0 for x, y in dataloader: x = x.reshape(-1, 1).float() y = y.reshape(-1, 1).float() pred = model(x).reshape(-1, 1) optimizer.zero_grad() loss = F.mse_loss(pred, y) # mu, sigma = model.mu_sigma(x) # dist = Normal(mu, sigma) # values = dist.log_prob(y) # loss += torch.sum(values) loss.backward() optimizer.step() running_loss += loss.item() print(running_loss / len(dataloader))
问题与疑问
现在调用model.mu_sigma(x)时,所有输出的mu均为0,sigma均为1。请问我哪里出现了错误?有哪些更合适的损失函数可以选择?
错误分析与解决
输出层ReLU导致的数值截断
mu的最后一层使用了ReLU,但目标动作包含负值(如-1),ReLU会将所有小于0的输出强制设为0,直接导致mu无法学习到负值,始终输出0。logvar的最后一层使用ReLU会限制其输出为非负值,网络初始化时线性层的权重偏置通常接近0,经过ReLU后logvar输出为0,计算得到sigma = exp(0.5*0) = 1,因此sigma始终固定为1。
解决方法:移除
mu和logvar层的ReLU激活函数:self.mu = nn.Linear(16, 1) # 直接输出均值,无激活 self.logvar = nn.Linear(16, 1) # logvar可以是任意实数,exp后保证方差为正训练逻辑的优化
当前使用采样动作和目标动作的MSE损失,但采样动作带有正态分布的随机性,会给损失带来额外噪声,不利于均值和方差的稳定学习。建议直接基于正态分布的参数计算损失。
合适的损失函数选择
负对数似然(NLL)损失:这是拟合正态分布到狄拉克δ分布的标准损失。对于目标动作
y,正态分布的对数概率为dist.log_prob(y),我们需要最小化负的对数概率均值:mu, sigma = model.mu_sigma(x) dist = torch.distributions.Normal(mu, sigma) loss = -dist.log_prob(y).mean()这个损失会直接驱动均值
mu逼近目标动作y,同时让方差sigma尽可能小,完美匹配用正态分布拟合狄拉克δ分布的需求。KL散度:拟合正态分布
q(a|x)到狄拉克δ分布p(a|x)=δ(a-y)时,KL散度KL(q||p)的计算等价于负对数似然(狄拉克分布的对数概率在a=y处为无穷大,KL散度的优化目标最终简化为最大化log q(y|x))。无需额外实现狄拉克分布模块,直接用NLL损失即可达到KL散度的优化效果。MSE损失:可作为辅助损失,但单独使用时对缩小方差的驱动较弱,因为它只关注采样动作和目标的差异,没有直接约束方差参数。如果使用,建议结合NLL损失一起优化。
内容的提问来源于stack exchange,提问作者Tigertron

