torch._C._infer_size()含义解析——基于MultivariateNormal的KL散度计算
理解PyTorch中torch._C._infer_size()在MultivariateNormal KL散度计算中的作用
核心功能
torch._C._infer_size(a_shape, b_shape)是PyTorch内部的工具函数,作用是根据PyTorch的广播规则,计算两个输入形状能兼容的统一广播形状——简单说就是找到两个形状按照广播逻辑扩展后能匹配的共同形状。
在MultivariateNormal KL散度中的具体作用
在多元正态分布的KL散度计算中,两个分布p和q可能带有不同的batch维度(比如p是单样本分布,q是批量分布)。_unbroadcasted_scale_tril的最后两个维度是协方差三角矩阵的尺寸(比如n×n),前面的所有维度都是batch维度,代码里的shape[:-2]就是提取这些batch维度的形状。
_infer_size在这里的任务:
- 提取
p和q的batch形状 - 推导它们能广播到的统一batch形状(即
combined_batch_shape) - 后续将
p和q的协方差三角矩阵都通过expand方法扩展到这个统一形状,确保后续矩阵求解、马氏距离计算等操作能在匹配的batch维度下进行,避免形状不兼容的报错。
示例场景
举个实际例子:
- 若
p._unbroadcasted_scale_tril.shape为(2, 3, 3),对应batch形状(2,),是3维多元正态分布 - 若
q._unbroadcasted_scale_tril.shape为(4, 2, 3, 3),对应batch形状(4,2)
经过_infer_size计算后,combined_batch_shape会得到(4,2),之后两个协方差矩阵都会被扩展为(4,2,3,3),这样每个batch位置的分布都能一一对应进行KL散度计算。
补充说明
这个内部函数的逻辑和PyTorch公开的torch.broadcast_shapes()完全一致,只是_infer_size是底层实现,早期版本的PyTorch KL散度计算用了这个内部函数,后续版本可能已替换为公开的broadcast_shapes。
内容的提问来源于stack exchange,提问作者dasen
相关产品推荐
相关产品推荐

