np.trace()专用于计算二维数组主对角线和,支持offset偏移与dtype精度控制;对高维数组需显式指定axis1/axis2,非方阵会报错,性能优于sum(diag())且不产生中间数组。

trace 函数直接计算二维数组主对角线元素和
numpy.trace() 是专为提取并求和主对角线(从左上到右下)设计的函数,只对二维数组(即 shape 为 (n, n) 或 (m, n))有效。它默认取 axis1=0, axis2=1 对应的两个轴,把其余维度视为批量处理维度。
常见错误是传入一维数组或三维以上数组却没指定轴——此时会报 ValueError: Input must be 2-D。若你有一个形状为 (5, 4, 4) 的批量矩阵,可以加参数 axis1=1, axis2=2 让它沿后两维算迹,结果为长度为 5 的一维数组。
- 对
arr = np.array([[1, 2], [3, 4]]),np.trace(arr)返回5 - 对
arr = np.arange(24).reshape(2, 3, 4),np.trace(arr, axis1=1, axis2=2)报错:因为3 != 4,非方阵无法取主对角线 - 若想强制取前 min(m,n) 个元素和,得手动切片:
np.diag(arr).sum()(但np.diag()对非方阵返回的是扁平化对角线,行为与trace不同)
trace 支持 offset 和 dtype 控制精度与范围
用 offset 参数可偏移对角线:正数向右上,负数向左下。比如 np.trace(arr, offset=1) 取次对角线(位置 (0,1), (1,2), ...),常用于计算上三角紧邻对角线的和。
dtype 参数影响累加过程的中间类型。例如对 uint8 数组大量求和时,不设 dtype=np.uint64 可能溢出;而浮点数组默认用 float64 累加,若需控制内存可显式设为 dtype=np.float32,但要注意精度损失。
-
np.trace(np.eye(3), offset=1)→0.0(3×3 单位阵的 offset=1 对角线只有两个位置,全为 0) -
np.trace(np.ones((1000, 1000), dtype=np.uint8))直接返回1000(没溢出),但换成np.ones((3000, 3000), dtype=np.uint8)就会回绕成3000 % 256 = 148 - 安全做法:
np.trace(big_uint8_arr, dtype=np.int64)
trace 与手写 sum(diag()) 性能差异不大,但语义更明确
有人习惯写 np.sum(np.diag(arr)),这在小数组上几乎没区别,但有两点隐患:一是 np.diag(arr) 会先构造一个长度为 min(m,n) 的新数组,浪费内存;二是当 arr 是视图(如切片)时,np.diag() 返回的是副本,而 np.trace() 是纯计算,不分配中间存储。
更重要的是语义——trace 明确表达“我要矩阵的迹”,而 sum(diag()) 更像临时拼凑的技巧,在代码审查或协作中容易引发疑问。
- 对
arr = np.random.rand(1000, 1000),np.trace(arr)比np.sum(np.diag(arr))快约 15%~20%,主要省去了diag的内存分配开销 - 若
arr是arr[::2, ::2]这样的跨步视图,np.diag(arr)会触发完整拷贝,np.trace()则无此问题
高维数组用 trace 要小心 axis 顺序和广播规则
当数组维度超过二维,np.trace() 默认只对前两个轴操作。比如 arr.shape == (2, 3, 4, 5),np.trace(arr) 等价于对每个 (3,4) 子矩阵沿第 0 和第 1 轴(即索引 0 和 1)取迹,结果 shape 是 (2, 5)。如果你本意是对每组 (4,5) 算迹,就必须显式写 np.trace(arr, axis1=2, axis2=3)。
另一个易错点:trace 不支持 broadcasting。如果传入两个不同 shape 的数组做运算后再 trace,得先确保它们已对齐,否则会得到意外的降维结果或报错。
-
arr = np.ones((2, 3, 3)); np.trace(arr).shape→(2,)(对每个(3,3)算一个标量) -
np.trace(arr, axis1=1, axis2=2).shape→ 同样是(2,),但含义一致 - 若误写
np.trace(arr, axis1=0, axis2=1),则对每个(2,3)子矩阵算迹,结果 shape 是(3,),逻辑完全偏离预期
np.einsum('ii->i', arr, axes=[axis1, axis2]) 或手动索引更稳妥。Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











