You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 12:20:40