如何仅计算NumPy matmul的实部?含einsum扩展需求
复数矩阵乘法取实部的高效实现及einsum推广
基础场景:矩阵乘法取实部
复数矩阵乘积的实部可以直接通过实部矩阵相乘减去虚部矩阵相乘得到,无需计算完整的复数矩阵乘积再取实部,这样能减少约一半的浮点乘法操作,大幅提升性能。
原理
对于复数 ( z_1 = a + bi ) 和 ( z_2 = c + di ),乘积的实部为 ( ac - bd )。推广到矩阵乘法,若 ( A = A_{\text{re}} + iA_{\text{im}} ),( B = B_{\text{re}} + iB_{\text{im}} ),则:
[
\text{Re}(A \times B) = A_{\text{re}} \times B_{\text{re}} - A_{\text{im}} \times B_{\text{im}}
]
代码实现
import numpy as np # 生成测试用复数数组 a = np.random.randn(1000, 2000).astype(np.complex128) + 1j * np.random.randn(1000, 2000).astype(np.complex128) b = np.random.randn(2000, 1500).astype(np.complex128) + 1j * np.random.randn(2000, 1500).astype(np.complex128) # 原方法 c_orig = np.matmul(a, b).real # 优化方法 c_opt = np.matmul(a.real, b.real) - np.matmul(a.imag, b.imag) # 验证结果一致性 assert np.allclose(c_orig, c_opt)
推广到einsum操作
对于多复数数组的einsum操作取实部,核心思路是拆分每个复数数组的实部和虚部,枚举所有取偶数个虚部的组合,根据虚部数量的奇偶性调整符号后求和,即可得到原复数einsum结果的实部。
原理
以4个复数数组的einsum "ab,bc,cd,da->a" 为例,复数乘积的实部由所有取0、2、4个虚部的组合项构成:
- 取0个虚部:所有数组用实部,符号为 ( (+1) )
- 取2个虚部:任意两个数组用虚部,其余用实部,每个组合符号为 ( (-1) )
- 取4个虚部:所有数组用虚部,符号为 ( (+1) )
将这些组合的einsum结果按符号相加,就是原复数einsum的实部。
代码实现
# 生成测试用复数数组 w = np.random.randn(500, 500).astype(np.complex128) + 1j * np.random.randn(500, 500).astype(np.complex128) x = np.random.randn(500, 500).astype(np.complex128) + 1j * np.random.randn(500, 500).astype(np.complex128) y = np.random.randn(500, 500).astype(np.complex128) + 1j * np.random.randn(500, 500).astype(np.complex128) z = np.random.randn(500, 500).astype(np.complex128) + 1j * np.random.randn(500, 500).astype(np.complex128) # 原方法 result_orig = np.einsum("ab,bc,cd,da->a", w, x, y, z).real # 拆分实部和虚部(numpy的real/imag返回视图,无额外内存开销) w_re, w_im = w.real, w.imag x_re, x_im = x.real, x.imag y_re, y_im = y.real, y.imag z_re, z_im = z.real, z.imag # 计算所有有效组合项 term0 = np.einsum("ab,bc,cd,da->a", w_re, x_re, y_re, z_re) term1 = -np.einsum("ab,bc,cd,da->a", w_re, x_re, y_im, z_im) term2 = -np.einsum("ab,bc,cd,da->a", w_re, x_im, y_re, z_im) term3 = -np.einsum("ab,bc,cd,da->a", w_re, x_im, y_im, z_re) term4 = -np.einsum("ab,bc,cd,da->a", w_im, x_re, y_re, z_im) term5 = -np.einsum("ab,bc,cd,da->a", w_im, x_re, y_im, z_re) term6 = -np.einsum("ab,bc,cd,da->a", w_im, x_im, y_re, z_re) term7 = np.einsum("ab,bc,cd,da->a", w_im, x_im, y_im, z_im) # 合并结果 result_opt = term0 + term1 + term2 + term3 + term4 + term5 + term6 + term7 # 验证结果一致性 assert np.allclose(result_orig, result_opt)
关键注意事项
- 内存效率:numpy中
array.real和array.imag返回的是原数组的视图,而非副本,不会额外占用大量内存。 - 性能提升:该方法将浮点乘法次数减少至原复数运算的一半,且实数组的einsum/matmul操作通常能获得numpy更优的底层优化,大数组场景下性能提升显著。
- 通用性:此方法可推广到任意数量复数数组的einsum操作,只需枚举所有偶数个虚部的组合并按符号调整即可。
内容的提问来源于stack exchange,提问作者nobe
相关产品推荐
相关产品推荐

