如何通过向量化方式从Numpy多元正态分布中批量采样?
Numpy向量化批量多元正态分布采样实现方案
批量多元正态采样可以通过多元分布的采样拆解逻辑实现完全向量化,无需显式循环,也不需要引入PyTorch做数据格式转换。
核心逻辑基于多元正态分布的采样公式:采样结果 = 均值 + 协方差Cholesky下三角矩阵 @ 标准正态采样向量,Numpy的线性代数算子原生支持批量维度输入,可以直接适配每个样本对应独立协方差矩阵的场景。
实现代码
假设输入批量均值a形状为*(k, N),批量协方差矩阵b形状为(k, N, N),输出采样结果c形状为(k, N)*,实现如下:
import numpy as np # 构造测试输入(实际使用时替换为你的a、b即可) k = 32 # 批量大小 N = 8 # 多元分布维度 a = np.random.randn(k, N) # 构造半正定协方差矩阵示例,你的输入b可跳过这一步,只要保证是合法半正定矩阵即可 b = np.random.randn(k, N, N) b = b @ b.transpose(0, 2, 1) + 1e-6 * np.eye(N)[np.newaxis, ...] # 向量化采样逻辑 # 1. 批量Cholesky分解,得到每个协方差对应的下三角矩阵,输出形状(k, N, N) L = np.linalg.cholesky(b) # 2. 生成批量标准正态采样,新增最后一维适配批量矩阵乘法 eps = np.random.randn(k, N, 1) # 3. 线性变换后降维得到最终采样结果 c = a + (L @ eps).squeeze(axis=-1)
验证说明
固定随机种子后,上述实现的输出和原循环版本的输出完全一致。相比Python层的显式循环,向量化实现的性能在批量规模较大时会有数量级的提升,且全程使用Numpy算子,无需额外的数据格式转换。
内容的提问来源于stack exchange,提问作者potpie
相关产品推荐
相关产品推荐

