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

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未定义错误。

修复步骤:

  1. 添加初始状态转移:指定从模型起始节点到具体状态的概率
  2. 添加状态间转移关系:即使是监督训练,也需要先定义基础转移结构,后续fit会用标注数据更新概率
  3. 调整执行顺序:完成状态、初始状态、转移关系添加后,再执行bake()和fit
  4. 修正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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.11 12:05:20