np.equal的dtype关键字参数为何未生效?
关于numpy中
np.equal的dtype参数未生效的问题 我完全懂你这会儿的困惑——明明numpy文档里标注dtype是np.equal的有效关键字参数,但实际测试时指定dtype=float,返回结果的类型居然还是bool,反观np.add这类算术型ufunc却能乖乖遵循dtype设置,这反差确实让人摸不着头脑。
先把你给出的测试代码再贴一遍,方便大家看清楚场景:
import numpy as np np.__version__ # 输出 '1.14.2' a = b = np.arange(2).astype(np.uint8) np.equal(a, b, dtype=float).dtype # 结果是 dtype('bool') np.add(a, b).dtype # 结果是 dtype('uint16')
背后的原因
这其实是numpy中ufunc的设计逻辑导致的:比较类ufunc(比如np.equal、np.greater、np.less这类)的输出类型是固定为布尔型的,它们的核心功能就是返回两个数组元素的比较结果,所以内部会直接忽略你传入的dtype参数。而算术类ufunc(比如np.add、np.subtract)则需要根据输入类型和指定的dtype来计算输出的数值类型,因此dtype参数能正常生效。
文档的小疏漏
之所以文档里会标注dtype是有效参数,大概率是因为numpy的很多ufunc共享了通用的参数定义模板,比较类ufunc也继承了这个模板,但实际实现时并没有用到dtype参数,属于文档的小疏漏。
替代解决方案
如果你确实需要把比较结果转换成指定类型,最简单的办法就是在比较完成后显式转换:
result = np.equal(a, b).astype(float) result.dtype # 此时输出就是 dtype('float64')
这样就能得到你想要的类型啦。
内容的提问来源于stack exchange,提问作者cnapun
相关产品推荐
相关产品推荐

