基于广播机制的Python逐元素乘法:独热向量筛选子矩阵
用NumPy广播实现独热向量与高维矩阵的匹配乘法
嘿,这个需求用NumPy的广播机制就能轻松搞定,核心思路是让独热向量的维度和目标矩阵的维度对齐,这样广播会自动帮我们完成逐元素的匹配乘法。我给你一步步拆解实现过程:
步骤1:明确输入形状
你的例子里:
- 独热向量
v可以是(1,3)的二维数组,也可以简化成(3,)的一维数组,后者更简洁 - 矩阵
m是(3,2,3)(对应n=3, m=2, r=3的n*m*r形状)
步骤2:调整独热向量的形状
要实现广播,我们需要把v的形状调整为(3,1,1)——这样它的第一个维度和m的第一个维度(n)对应,后面两个维度是单元素,会被广播到m的m和r维度上。
你可以用两种方式调整形状:
方式A:用reshape
import numpy as np # 定义输入 v = np.array([0.0, 1.0, 0.0]) # 一维数组版本 m = np.array([[[1,2,3],[4,5,6]], [[5,6,7],[7,8,9]], [[2,4,7],[1,8,9]]]) # 调整v的形状为(n,1,1) v_aligned = v.reshape(-1, 1, 1)
方式B:用np.newaxis增加维度
v_aligned = v[:, np.newaxis, np.newaxis]
如果你的v是(1,3)的二维数组(比如v = np.array([[0.0,1.0,0.0]])),只需要先转置再增加维度:
v_aligned = v.T[:, np.newaxis]
步骤3:执行逐元素乘法
直接用*运算符做逐元素乘法,广播会自动把v_aligned扩展成和m一样的(3,2,3)形状,完成匹配:
prod = m * v_aligned
验证结果
打印prod就能得到你想要的输出:
print(prod)
输出:
[[[0 0 0] [0 0 0]] [[5 6 7] [7 8 9]] [[0 0 0] [0 0 0]]]
原理说明
广播机制会自动将v_aligned中每个元素(0或1)复制到对应子矩阵的所有位置,然后和m的对应元素相乘。因为独热向量里只有一个1,所以只有对应的子矩阵会被保留,其余子矩阵全部置零。
内容的提问来源于stack exchange,提问作者Mahmood
相关产品推荐
相关产品推荐

