如何更高效计算马尔可夫链的分布统计量?
问题
我有一个表示齐次马尔可夫链的2000×2000维度概率转移矩阵,想要获取该链前200步中每个概率分布的统计量(每一步的首行分布),已编写如下代码:
using Distributions, LinearAlgebra # 定义转移矩阵 function tm(N::Int, n0::Int) [pdf(Hypergeometric(N-l,l,n0),k-l) for l in 0:N, k in 0:N] end # 计算概率向量的5分位数 function percentile5(M::Vector) s=0 i=0 while s <= 0.05 i += 1 s += M[i] end return i-1 end # 计算统计量矩阵:行是均值、5分位数、标准差;列是每一步 function stats(N::Int, n0::Int, m::Int) A = tm(N,n0) B = I # 初始化为单位矩阵 sup = 0:N # 分布的支撑集 sup2 = [k^2 for k in sup] stats = zeros(3,m) for i in 1:m C = B[1,:] stats[1,i] = sum(C .* sup) # 均值 stats[2,i] = percentile5(C) # 5分位数 stats[3,i] = sqrt(sum(C .* sup2) - stats[1,i]^2) # 标准差 B = A*B end return stats end data = stats(2000,50,200)
请问是否存在更高效(更快)的方法完成相同计算?目前我未想到更好方案,想了解可提速的技巧。
优化提速方案
1. 放弃全矩阵维护,只跟踪首行分布
原代码中每次执行B = A*B是O(N³)的矩阵乘法,但我们只需要B的首行数据。改为维护单个一维分布向量,每次更新为向量与转移矩阵的乘积(O(N²)操作),计算量直接降低2000倍(N=2000时)。
修改后的核心逻辑:
function stats(N::Int, n0::Int, m::Int) A = tm(N,n0) current_dist = zeros(N+1) current_dist[1] = 1.0 # 初始分布对应状态0(1-based索引) sup = 0:N sup2 = sup.^2 stats_mat = zeros(3,m) for i in 1:m stats_mat[1,i] = dot(current_dist, sup) # 用BLAS优化的dot替代sum(C.*sup) stats_mat[2,i] = percentile5(current_dist) stats_mat[3,i] = sqrt(dot(current_dist, sup2) - stats_mat[1,i]^2) current_dist = current_dist * A # 仅更新当前分布向量 end return stats_mat end
2. 用稀疏矩阵存储转移矩阵
观察转移矩阵结构:当k < l或k > l + n0时,概率值为0,矩阵大部分元素是无效的0值。用稀疏矩阵存储可大幅减少内存占用和乘法计算量。
修改转移矩阵构造函数:
using SparseArrays function tm(N::Int, n0::Int) rows = Int[] cols = Int[] vals = Float64[] for l in 0:N # 计算k的有效范围(仅存储非零概率) min_k = max(l, n0) max_k = min(N, l + n0) for k in min_k:max_k p = pdf(Hypergeometric(N-l, l, n0), k - l) p > 0 && (push!(rows, l+1); push!(cols, k+1); push!(vals, p)) end end return sparse(rows, cols, vals) end
3. 优化百分位数计算
原循环累加的方式可改用二分查找加速,利用累积和的单调性将时间复杂度从O(N)降到O(logN):
function percentile5(M::Vector) cum = cumsum(M) # 找到第一个超过0.05的位置,返回前一个索引(对应状态值) i = searchsortedfirst(cum, 0.05 + eps()) return isnothing(i) ? length(M)-1 : i-1 end
4. 预分配内存减少复制
使用mul!函数在预分配的向量中存储更新后的分布,避免每次循环重新分配内存:
function stats(N::Int, n0::Int, m::Int) A = tm(N,n0) current_dist = zeros(N+1) current_dist[1] = 1.0 next_dist = similar(current_dist) # 预分配更新后的分布向量 sup = 0:N sup2 = sup.^2 stats_mat = zeros(3,m) for i in 1:m stats_mat[1,i] = dot(current_dist, sup) stats_mat[2,i] = percentile5(current_dist) stats_mat[3,i] = sqrt(dot(current_dist, sup2) - stats_mat[1,i]^2) mul!(next_dist, current_dist, A) # 直接写入预分配向量 current_dist, next_dist = next_dist, current_dist # 交换引用避免复制 end return stats_mat end
5. 利用分布性质递推均值和方差
均值和方差无需依赖完整分布,可通过递推公式直接计算,将这部分计算从O(N)降到O(1):
function stats(N::Int, n0::Int, m::Int) A = tm(N,n0) current_dist = zeros(N+1) current_dist[1] = 1.0 next_dist = similar(current_dist) stats_mat = zeros(3,m) mu = 0.0 sigma2 = 0.0 for i in 1:m # 百分位数仍需依赖分布 stats_mat[2,i] = percentile5(current_dist) # 递推均值和方差 if i == 1 mu = 0.0 sigma2 = 0.0 else # 均值递推公式 mu = mu*(1 - n0/N) + n0 # 方差递推公式 term1 = sigma2 term2 = (n0*(N - n0))/(N^2*(N-1)) * (N*mu - (sigma2 + mu^2)) term3 = 2*(n0/N)*(N*mu - (sigma2 + mu^2)) sigma2 = term1 + term2 + term3 end stats_mat[1,i] = mu stats_mat[3,i] = sqrt(sigma2) mul!(next_dist, current_dist, A) current_dist, next_dist = next_dist, current_dist end return stats_mat end
内容的提问来源于stack exchange,提问作者user4901852
相关产品推荐
相关产品推荐

