Pomegranate训练HMM监督学习遇UnboundLocalError及转移矩阵查看问题
问题描述
- 使用Pomegranate 0.14.4训练监督学习隐马尔可夫模型(HMM),目标是基于观测数据预测状态
- 运行代码时在
model.bake()步骤报错:UnboundLocalError: local variable 'dist' referenced before assignment - 需了解如何查看模型的转移矩阵
用户代码:
from pomegranate import * import numpy as np # Supervised method that calculates the transition matrix: d1 = State(UniformDistribution.from_samples([3.243221498397177, 3.210684537495482, 3.227662201472816, 3.286410817416738, 3.290573650708864, 3.286058136226862, 3.266480693857006])) d2 = State(UniformDistribution.from_samples([3.449282367485096, 1.97317859465635, 1.897551432353011, 3.454609351559659, 3.127357456033111, 1.779308337786426, 3.802891929694426, 3.359766157565077, 2.959428499979418])) d3 = State(UniformDistribution.from_samples([1.892812118441474, 1.589353118681066, 2.09269978285637, 2.104391496570218, 1.656771181054144])) model = HiddenMarkovModel() model.add_states(d1, d2, d3) # print(model.to_json()) model.bake() model.fit([3.2, 6.7, 10.55], labels=[1, 2, 3], algorithm='labeled') all_pred = model.predict([2.33, 1.22, 1.4, 10.6])
错误详情:
File "C:\Program Files\JetBrains\PyCharm Community Edition 2021.2\plugins\python-ce\helpers\pydev\_pydev_bundle\pydev_umd.py", line 198, in runfile pydev_imports.execfile(filename, global_vars, local_vars) # execute the script File "C:\Program Files\JetBrains\PyCharm Community Edition 2021.2\plugins\python-ce\helpers\pydev\_pydev_imps\_pydev_execfile.py", line 18, in execfile exec(compile(contents+"\n", file, 'exec'), glob, loc) File "C:/Users/", line 774, in <module> model.bake() File "pomegranate/hmm.pyx", line 1047, in pomegranate.hmm.HiddenMarkovModel.bake UnboundLocalError: local variable 'dist' referenced before assignment
报错原因与修复
报错核心是HMM未定义初始状态和状态转移关系就执行了bake()。Pomegranate的HMM在bake()时需要明确初始状态概率及状态间转移逻辑,否则内部初始化会找不到必要的分布变量,触发dist未定义错误。
修复步骤:
- 添加初始状态转移:指定从模型起始节点到具体状态的概率
- 添加状态间转移关系:即使是监督训练,也需要先定义基础转移结构,后续
fit会用标注数据更新概率 - 调整执行顺序:完成状态、初始状态、转移关系添加后,再执行
bake()和fit - 修正
fit输入格式:观测数据需为序列列表(单条序列也要用双层列表包裹)
修复后示例代码:
from pomegranate import * import numpy as np # 定义状态与观测分布 d1 = State(UniformDistribution.from_samples([3.243221498397177, 3.210684537495482, 3.227662201472816, 3.286410817416738, 3.290573650708864, 3.286058136226862, 3.266480693857006])) d2 = State(UniformDistribution.from_samples([3.449282367485096, 1.97317859465635, 1.897551432353011, 3.454609351559659, 3.127357456033111, 1.779308337786426, 3.802891929694426, 3.359766157565077, 2.959428499979418])) d3 = State(UniformDistribution.from_samples([1.892812118441474, 1.589353118681066, 2.09269978285637, 2.104391496570218, 1.656771181054144])) model = HiddenMarkovModel() model.add_states(d1, d2, d3) # 添加初始状态转移(从start节点到d1,概率1.0) model.add_transition(model.start, d1, 1.0) # 添加状态间转移示例(后续fit会用标注数据更新) model.add_transition(d1, d2, 0.3) model.add_transition(d2, d3, 0.3) model.add_transition(d3, d1, 0.3) # 添加到end节点的转移 model.add_transition(d1, model.end, 0.1) model.add_transition(d2, model.end, 0.1) model.add_transition(d3, model.end, 0.1) # 先bake再执行fit model.bake() # fit输入需为双层列表,标签也对应双层结构 model.fit([[3.2, 6.7, 10.55]], labels=[[1, 2, 3]], algorithm='labeled') all_pred = model.predict([2.33, 1.22, 1.4, 10.6]) print(all_pred)
查看转移矩阵的方法
训练完成后,可通过两种方式获取转移矩阵:
1. 遍历转移关系直接查看
# 遍历所有转移规则,输出状态对与对应概率 for transition in model.transitions: start = transition.state.name if hasattr(transition.state, 'name') else str(transition.state) end = transition.next_state.name if hasattr(transition.next_state, 'name') else str(transition.next_state) print(f"从{start}到{end}的转移概率:{transition.probability}")
2. 构建结构化矩阵
先给状态命名,再生成矩阵形式的转移表:
# 给状态设置名称,方便识别 d1.name = "State1" d2.name = "State2" d3.name = "State3" # 过滤掉start和end节点,仅保留自定义状态 states = [s for s in model.states if s not in [model.start, model.end]] state_names = [s.name for s in states] state_count = len(states) # 初始化转移矩阵 transition_matrix = np.zeros((state_count, state_count)) # 填充矩阵数据 for i, start_state in enumerate(states): for j, end_state in enumerate(states): for trans in model.transitions: if trans.state == start_state and trans.next_state == end_state: transition_matrix[i][j] = trans.probability break # 打印结果 print("转移矩阵:") print(transition_matrix) print("状态顺序:", state_names)
内容的提问来源于stack exchange,提问作者maximus
相关产品推荐
相关产品推荐

