使用PyMC3进行MCMC光谱解混时遭遇KeyError: 'abundances'问题
光谱解混MCMC代码报错:KeyError: 'abundances'
尝试用MCMC算法实现光谱解混,运行代码时遇到KeyError: 'abundances',切换NUTS和Metropolis采样器问题依旧。
运行的代码片段
import pymc3 as pm import theano.tensor as tt import arviz as az import numpy as np # 原代码遗漏的必要导入 """ Perform the spectral unmixing through a MCMC algorithm """ ### SHOW AVAILABLE END MEMBERS for l in range(len(TOT_names)): print('Index: '+str(l)+', Compound: '+str(TOT_names[l])) ### ENDMEMBERS my_arr = [] lst = list(map(int, input("Which end-members (enter comma separated values): ").split(","))) for k in lst: my_arr.append(EM[k][i:j]) ENDM = pm.floatX(np.array(my_arr)) print('ENDMEMBERS ARRAY: ', ENDM.shape) ### DATA MATRIX MATRIX = JM0340_matrix ### DEFINE MCMC MODEL with pm.Model() as model: # Prior distributions for the abundances ABUNDANCES = pm.Dirichlet('abundances', a=np.ones(len(lst))) #print('ABUNDANCES ARRAY: ', ABUNDANCES) # Constraint on non-negativity of abundances pm.Potential('AB_POS_CONSTRAINT', pm.math.switch(pm.math.sum(pm.math.maximum(ABUNDANCES, 0)) - pm.math.sum(ABUNDANCES) < 0, -np.inf, 0)) # Constraint on the sum of abundances AB_SUM = pm.Deterministic('ab_sum', pm.math.sum(ABUNDANCES)) pm.Potential('AB_SUM_CONSTRAINT', pm.math.switch(tt.abs_(AB_SUM - 1) > 1e-3, -np.inf, 0)) # Compute modeled spectra MODELED_SPECTRA = pm.Deterministic('modeled_spectra', pm.math.dot(ABUNDANCES, ENDM)) # Likelihood distribution for the data DATA = pm.Normal('data', mu=MODELED_SPECTRA, sd=1, observed=MATRIX) #sd standard deviation # Sampling Define MCMC model TRACE = pm.sample(draws=400, tune=210, chains=8, cores=8, step=pm.Metropolis(), return_inferencedata=True) # Extract the trace of the abundances ABUND = TRACE['abundances'].mean(axis=0) print('ABUNDANCES ARRAY: ', ABUND.shape)
报错信息
81 # Extract the trace of the abundances ---> 82 ABUND = TRACE['abundances'].mean(axis=0) 83 print('ABUNDANCES ARRAY: ', ABUND.shape) File ~\anaconda3\lib\site-packages\arviz\data\inference_data.py:236, in InferenceData.__getitem__(self, key) 234 """Get item by key.""" 235 if key not in self._groups_all: --> 236 raise KeyError(key) 237 return getattr(self, key) KeyError: 'abundances'
解决方案
修正变量提取方式
当pm.sample设置return_inferencedata=True时,返回的是ArviZ的InferenceData对象,变量并非直接存储在根对象中,而是放在posterior组里。正确的提取方式如下:# 提取abundances并计算均值 ABUND = TRACE.posterior['abundances'].mean(dim=('chain', 'draw')).values print('ABUNDANCES ARRAY: ', ABUND.shape)这里用
dim=('chain', 'draw')指定对链和采样维度求均值,最后用.values转换为numpy数组。移除多余的约束
Dirichlet分布本身就满足非负性和和为1的约束,代码中添加的AB_POS_CONSTRAINT和AB_SUM_CONSTRAINT完全多余,甚至可能干扰采样过程,建议直接删除这两段代码。补充缺失的依赖导入
原代码中使用了np.array但没有导入numpy,需在开头添加import numpy as np。
内容的提问来源于stack exchange,提问作者Alex
相关产品推荐
相关产品推荐

