You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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
---&gt; 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:
--&gt; 236     raise KeyError(key)
    237 return getattr(self, key)

KeyError: 'abundances'

解决方案

  1. 修正变量提取方式
    当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数组。

  2. 移除多余的约束
    Dirichlet分布本身就满足非负性和和为1的约束,代码中添加的AB_POS_CONSTRAINT和AB_SUM_CONSTRAINT完全多余,甚至可能干扰采样过程,建议直接删除这两段代码。

  3. 补充缺失的依赖导入
    原代码中使用了np.array但没有导入numpy,需在开头添加import numpy as np。

内容的提问来源于stack exchange,提问作者Alex

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.23 13:18:19