Numpy:实现二维数组与三维数组最后维度切片的逐元素乘法
Numpy:实现二维数组与三维数组最后维度切片的逐元素乘法
嘿,这个问题我之前也踩过坑!Numpy的维度广播规则有时候确实容易让人困惑,不过解决起来其实挺简单的~
你直接用x * y不行的原因是维度不匹配:x是二维数组(100,10),而y是三维数组(100,10,2)。Numpy的广播需要从最后一个维度开始对齐,这里x的最后一个维度是10,y的倒数第二个维度是10、倒数第一个是2,没法直接对齐,所以会报错。
要实现你想要的“每个y[i,j,:]乘以x[i,j]”,关键是给x增加一个长度为1的新维度,让它的形状变成(100,10,1),这样就能和y的(100,10,2)触发广播机制了。下面给你几种常用的实现方式:
方法1:用[..., None]扩展维度(最简洁)
这是Numpy里最常用的扩展维度写法,...表示取前面所有维度,None等价于np.newaxis,用来新增维度:
import numpy as np x = np.random.randn(100, 10) y = np.random.randn(100, 10, 2) result = x[..., None] * y
方法2:用np.expand_dims显式扩展维度
如果觉得上面的写法不够直观,可以用np.expand_dims指定要扩展的轴(axis=-1表示最后一个轴):
result = np.expand_dims(x, axis=-1) * y
方法3:用reshape手动修改形状
你也可以手动修改x的形状,在末尾加一个1:
result = x.reshape(x.shape + (1,)) * y
验证结果是否正确
你可以随便取一组i,j的值手动计算,验证结果是否符合预期:
# 取第0行第0列的元素验证 manual_calc = x[0, 0] * y[0, 0, :] print(np.allclose(result[0, 0, :], manual_calc)) # 输出True说明结果正确
本质上这几种方法都是同一个思路:让x的维度和y的前两个维度对齐,同时新增一个长度为1的维度,让Numpy自动把这个维度广播到和y的最后一个维度匹配,这样每个x[i,j]就会和y[i,j,0]、y[i,j,1]分别相乘啦~
备注:内容来源于stack exchange,提问作者Euler_Salter
相关产品推荐
相关产品推荐

