Cupy中cp.transpose(cp.nonzero(a))报错问题求助
解决CuPy中
cp.transpose(cp.nonzero(a))报错的问题 问题根源
你遇到的报错是因为CuPy的transpose函数只能作用于CuPy数组对象,而cp.nonzero(a)返回的是一个由CuPy数组组成的tuple(和NumPy的返回格式一致)。NumPy的transpose在这里做了隐式处理:自动将tuple转换为NumPy数组后再执行转置,但CuPy没有这个隐式转换逻辑,所以直接传入tuple会触发AttributeError。
两种可行的解决方法
方法1:显式将
tuple转为CuPy数组后再转置
先把cp.nonzero(a)的结果转换成CuPy数组,再调用transpose,和你原来的NumPy写法逻辑完全对齐:cp.transpose(cp.array(cp.nonzero(a)))方法2:直接使用
cp.argwhere()(更简洁高效)
CuPy的argwhere函数可以直接返回非零元素的索引数组,形状为(非零元素数量, 数组维度),和np.transpose(np.nonzero(a))的输出完全一致,不需要手动处理转置:cp.argwhere(a)
验证示例
比如对于2D数组:
import cupy as cp a = cp.array([[1, 0, 3], [0, 5, 0]]) # 方法1输出 print(cp.transpose(cp.array(cp.nonzero(a)))) # 方法2输出 print(cp.argwhere(a))
两者都会输出:
[[0 0] [0 2] [1 1]]
内容的提问来源于stack exchange,提问作者Fernando Chueca
相关产品推荐
相关产品推荐

