如何将numpy张量指定维度的对角线元素设置为1?
实现方案
原代码失效原因
np.fill_diagonal 属于原地修改数组的方法,没有有效返回值(返回None),而np.apply_along_axis依赖回调函数的返回值构建结果,因此无法完成对原数组的修改。
最优实现方式
直接使用NumPy高级索引即可一行完成需求:
q[:, np.arange(3), np.arange(3)] = 1
索引逻辑说明
- 第一个维度的
:选中0轴上全部的3×3子矩阵 - 后两个维度同步传入
np.arange(3),精准匹配所有j=k的对角线坐标位
如果需要适配任意边长的子矩阵,可改用通用写法:
side_len = q.shape[1] q[:, np.arange(side_len), np.arange(side_len)] = 1
内容的提问来源于stack exchange,提问作者cjm2671
相关产品推荐
相关产品推荐

