Matlab中如何创建同时接收新输入和历史调用输出参数的函数句柄
Matlab带状态log后验函数句柄创建方案
核心思路
使用Matlab嵌套函数的变量捕获特性封装历史状态,生成仅接收theta作为单输入的函数句柄,完全符合hmcSampler的接口要求,历史的old_theta、old_logpdf由嵌套函数自行维护,无需第三方调用方传入。
完整实现步骤
1. 编写状态化后验函数生成工具
将状态维护逻辑、后验计算逻辑封装在同一个工厂函数中,返回符合接口要求的函数句柄:
function logpdf_handle = createStatefulLogPosterior(X, y, means, standevs, init_theta) % 初始化首次调用所需的历史参数 old_theta = init_theta; % 预计算初始old_logpdf intercept = init_theta(1); beta = init_theta(2:end); y_computed = X*beta + intercept; log_likelihood = log(y_computed); log_prior_params = 0; for i = 1:length(init_theta) lp = normalDistGrad(init_theta(i), means(i), standevs(i)); log_prior_params = log_prior_params + lp; end old_logpdf = log_likelihood + log_prior_params; % 嵌套函数:自动捕获父函数变量,状态持久化存储 function [logpdf, grad_logpdf] = statefulLogPosterior(theta) intercept = theta(1); beta = theta(2:end); y_computed = X*beta + intercept; log_likelihood = log(y_computed); del_loglikelihood = log_likelihood - old_logpdf; del_params = theta - old_theta; grad_params1 = del_loglikelihood ./ del_params; % 计算先验项与梯度 log_prior_params = 0; grad_params2 = []; for i = 1:length(theta) [lp, grad] = normalDistGrad(theta(i), means(i), standevs(i)); log_prior_params = log_prior_params + lp; grad_params2 = [grad_params2; grad]; end % 返回结果前更新历史状态 logpdf = log_likelihood + log_prior_params; grad_logpdf = grad_params1 + grad_params2; old_theta = theta; old_logpdf = logpdf; end % 返回仅接收单输入theta的函数句柄 logpdf_handle = @statefulLogPosterior; end % 正态分布梯度计算辅助函数 function [lpdf,glpdf] = normalDistGrad(X, Mu, Sigma) Z = (X - Mu)./Sigma; lpdf = sum(-log(Sigma) - .5*log(2*pi) - .5*(Z.^2)); glpdf = -Z./Sigma; end
2. 主脚本调用示例
替换你原有代码中创建函数句柄的部分即可:
%% Toy implementation of hmcsampler class in Matlab NumPredictors = 2; trueIntercept = 2; trueBeta = [3;0]; NumData = 100; rng('default') %For reproducibility X = rand(NumData,NumPredictors); mu = X*trueBeta + trueIntercept; y = mu; % define the mean and variance of normal distribution of each parameter means = [0; 0]; standevs = [1;1]; % create the startpoint from which sampling starts startpoint = randn(2, 1); % ************************* % 替换原有句柄创建代码 logpdf = createStatefulLogPosterior(X, y, means, standevs, startpoint); % ************************* % create an HMC sampler object smp = hmcSampler(logpdf, startpoint); % estimate maximum of log probability density [xhat, fitinfo] = estimateMAP(smp); num_chains = 4; chains = cell(num_chains, 1); burnin = 50000; num_samples = 2000000;
注意事项
- 每个
createStatefulLogPosterior生成的句柄状态独立,多链采样时每条链单独生成一个句柄即可避免状态冲突。 - 原代码修复了两处问题:
- 固定了遍历参数时写死
i=1:3的越界bug,改为动态读取theta长度适配参数维度 - 将矩阵除法
/改为元素级除法./,符合梯度计算的常规需求
- 固定了遍历参数时写死
内容的提问来源于stack exchange,提问作者shunyo
相关产品推荐
相关产品推荐

