使用meshgrid构建热力图时,如何解决np.dot维度不匹配问题?
解决热力图函数计算的维度不匹配问题
问题背景
原本构建Z = f(X,Y)热力图的代码可正常运行:
import numpy as np import seaborn as sb x = np.linspace(-np.pi, np.pi, 100) y = np.linspace(-np.pi, np.pi, 100) X,Y = np.meshgrid(x,y) f_XY = np.cos(X**2 + Y**2) ax = sb.heatmap(f_XY)
X和Y维度均为(100, 100),可逐对计算得到f_XY。
现在需要实现函数:计算每个点[x,y,2π]与向量[1,1,1]的点积的余弦值,数学表达式为f_XY = cos( dot([x,y,2pi], [1,1,1]) )。但运行以下代码时出现维度不匹配错误:
f_XY = np.cos( np.dot([X, Y, 2*np.pi], [1,1,1]) )
错误信息:
ValueError: shapes (3,200,200,1) and (3,) not aligned: 1 (dim 3) != 3 (dim 0)
解决方案
问题出在np.dot的使用方式上:将[X,Y,2π]传入np.dot时,numpy会把它们堆叠成高维数组,导致维度无法和[1,1,1]进行点积运算。
实际上,两个向量的点积就是对应元素相乘后求和,对于[x,y,2π]和[1,1,1],点积结果就是x*1 + y*1 + 2π*1,直接用逐元素运算即可:
import numpy as np import seaborn as sb x = np.linspace(-np.pi, np.pi, 100) y = np.linspace(-np.pi, np.pi, 100) X,Y = np.meshgrid(x,y) # 直接计算点积后取余弦 f_XY = np.cos(X + Y + 2 * np.pi) ax = sb.heatmap(f_XY)
如果后续需要更换权重向量(比如[a,b,c]),只需要对应修改为:
f_XY = np.cos(a * X + b * Y + c * 2 * np.pi)
内容的提问来源于stack exchange,提问作者WaterDrop
相关产品推荐
相关产品推荐

