Matlab中Householder三对角化的矩阵乘法为何效率极低?
问题背景
我正在实现一个Householder三对角化程序,处理n<10的小尺寸Hermitian矩阵,但在15万次调用时,核心矩阵更新步骤耗时约5秒,远慢于Matlab内置的hess.m(仅需零点几秒)。
现有代码
测试用Hermitian矩阵:
H = 1.0e-10 * [ 0.1386 + 0.0000i 0.0974 - 0.0260i 0.0434 + 0.0094i 0.0722 + 0.0670i 0.1128 + 0.1269i; 0.0974 + 0.0260i 0.0751 + 0.0000i 0.0288 + 0.0149i 0.0388 + 0.0616i 0.0557 + 0.1112i; 0.0434 - 0.0094i 0.0288 - 0.0149i 0.0146 + 0.0000i 0.0274 + 0.0164i 0.0444 + 0.0323i; 0.0722 - 0.0670i 0.0388 - 0.0616i 0.0274 - 0.0164i 0.0719 + 0.0000i 0.1216 + 0.0116i; 0.1128 - 0.1269i 0.0557 - 0.1112i 0.0444 - 0.0323i 0.1216 - 0.0116i 0.2105 + 0.0000i ];
三对角化函数:
function [T, U] = tridi(H) % Calculates a unitarily equivalent matrix T and unitary matrix U for a % given Hermitian matrix H, such that U'*H*U = T and T is of tridiagonal % form. % Setting size and starting point of iteration n = size(H,1); T = H; U = eye(n); % Loop over the first n-2 columns of H for k = 1 : n-1 % Householdertransformation on column-vector H(k+1:n,k) x = T(k+1:n,k); % Calculating secondary diagonal entry normValue = norm(x); % Initializing e1 e1 = zeros(length(x),1); e1(1) = 1; % phase of x1 phase = sign(x(1)); % Calculates normalized Householder-Vektor u u = x + norm(x)*phase*e1; u = u / norm(u); % Updating T and U: T(k+1:n,k+1:n) = (eye(n-k) - 2*(u*u'))*T(k+1:n,k+1:n)*(eye(n-k) - 2*(u*u')); U(2:n, k+1:n) = -phase*(U(2:n, k+1:n) - 2 * (U(2:n, k+1:n) * u) * u'); % Setting secondary diagonal entry of T T(k+1,k) = normValue; T(k,k+1) = normValue; % Setting appropriate row and column of T to zero T(k+2:n,k) = 0; T(k,k+2:n) = 0; end % Ensure that T is real T = real(T); end
瓶颈分析
核心瓶颈是这行矩阵更新代码:
T(k+1:n,k+1:n) = (eye(n-k) - 2*(u*u'))*T(k+1:n,k+1:n)*(eye(n-k) - 2*(u*u'));
15万次调用时这部分耗时约5秒。我尝试过将矩阵-矩阵乘法拆分为隐式的矩阵-向量乘法更新:
% Updating T and U: % Calculating P*T implicitly T(k+1:n, k+1:n) = T(k+1:n, k+1:n) - 2 * u * (u' * T(k+1:n, k+1:n)); % Calculating T*P and U*P implicitly T(k+1:n, k+1:n) = T(k+1:n, k+1:n) - 2 * (T(k+1:n, k+1:n) * u) * u'; U(2:n, k+1:n) = -phase*(U(2:n, k+1:n) - 2 * (U(2:n, k+1:n) * u) * u');
但没有性能提升,推测Matlab的内置矩阵乘法已经高度优化。
优化思路与困惑
我想进一步降低运行时间,目前考虑的方向:
- 利用T的Hermitian特性:每次迭代后T保持Hermitian,是否可以只计算下三角部分,再复制到上三角,避免重复计算?
- 用主对角线数组
alf和次对角线数组bet存储T,不过因为矩阵尺寸小,内存不是问题,不确定是否能提升速度; - 对于Hermitian矩阵的乘积(已知A、B对称可交换,乘积为Hermitian),不想用
A*B重复计算上下三角,熟悉tril和triu但担心仍有无效计算; - 纠结是否用循环单独计算下三角元素再复制——Matlab文档强调向量化,但小尺寸矩阵下循环会不会反而更快?
优化建议
针对小尺寸矩阵(n<10)的场景,给出几个针对性的优化方案:
1. 用Hermitian矩阵的简化公式更新T
Householder变换P = I - 2uu',对于Hermitian矩阵T,P*T*P可以简化为:P*T*P = T - 2u(u'*T) - 2(Tu)u' + 4u(u'*T u)u'
这个公式只需要计算向量和标量,避免了大临时矩阵的创建,效率更高:
uT = u' * T(k+1:n,k+1:n); Tu = T(k+1:n,k+1:n) * u; uTu = u' * Tu; T(k+1:n,k+1:n) = T(k+1:n,k+1:n) - 2*u*uT - 2*Tu*u' + 4*u*uTu*u';
2. 充分利用Hermitian特性减少冗余计算
更新T子矩阵时,只计算下三角(含对角线),再复制共轭转置到上三角,避免重复计算对称元素:
% 先计算完整的PT*P结果 uT = u' * T(k+1:n,k+1:n); Tu = T(k+1:n,k+1:n) * u; uTu = u' * Tu; T_sub = T(k+1:n,k+1:n) - 2*u*uT - 2*Tu*u' + 4*u*uTu*u'; % 保留下三角,复制共轭转置到上三角 T(k+1:n,k+1:n) = tril(T_sub) + triu(conj(T_sub'),1);
这能减少约一半的浮点运算量,对小矩阵的性能提升明显。
3. 小矩阵下尝试循环实现
Matlab的JIT编译器对小循环优化效果很好,对于n<10的矩阵,循环计算下三角元素可能比向量化矩阵乘法更快:
m = n - k; T_sub = T(k+1:n,k+1:n); uT = u' * T_sub; Tu = T_sub * u; uTu = u' * Tu; % 计算下三角元素 for i = 1:m for j = 1:i T_sub(i,j) = T_sub(i,j) - 2*u(i)*uT(j) - 2*Tu(i)*conj(u(j)) + 4*u(i)*uTu*conj(u(j)); end end % 复制到上三角 for i = 1:m for j = i+1:m T_sub(i,j) = conj(T_sub(j,i)); end end T(k+1:n,k+1:n) = T_sub;
4. 简化U矩阵的更新逻辑
原代码中U的索引和相位操作可以简化,直接用向量运算实现Householder变换:
U_col = U(:, k+1:n); U_col = U_col - 2*(U_col*u)*u'; U(:, k+1:n) = U_col;
可以把相位整合到u的计算中,避免后续额外的符号操作。
5. 参考内置函数的底层实现
Matlab内置的hess函数调用的是BLAS/LAPACK的zhetrd(复Hermitian)或dsytrd(实对称)函数,这些是高度优化的底层实现。如果必须自己实现,可以参考这些函数的逻辑,比如直接调用LAPACK接口:
% 直接调用LAPACK的zhetrd函数(复Hermitian矩阵三对角化) [alphas, betas, U] = zhetrd('V', H); % 构建三对角矩阵T T = diag(alphas) + diag(betas(1:end-1),1) + diag(conj(betas(1:end-1)),-1);
这几乎能达到内置函数的性能。
总结
对于n<10的小矩阵,核心优化点是利用Hermitian特性减少冗余计算,以及用向量运算替代矩阵乘法减少临时矩阵开销。可以优先尝试简化的P*T*P公式,再结合下三角计算+复制的方式,最后测试循环实现的性能——小矩阵下循环不一定比向量化慢。
内容的提问来源于stack exchange,提问作者Zedssad

