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

基于转移矩阵用NumPy向右填充路径数组:解决维度错误与向量化

马尔可夫链路径模拟错误解决与向量化实现

问题描述

给定状态与转移矩阵:

import numpy as np

n_states = 3
states = np.arange(n_states)

T = np.array([
    [0.5, 0.5, 0],
    [0.5, 0,  0.5],
    [0,  0,  1]
])

需要模拟n_sims条长度为n_steps的马尔可夫链路径,所有路径从状态0开始。初始化代码如下:

n_sims = 100
n_steps = 10
paths = np.zeros((n_sims, n_steps),  dtype=int)

尝试用np.random.Generator.choice填充路径时,编写了如下代码:

rng = np.random.default_rng(seed=123)

for s in range(1, n_steps+1):
    paths[:,s] = rng.choice(
        a=n_states,
        size=n_sim,
        p=T[paths[:,s-1]]
    )

触发错误:

ValueError: p must be 1-dimensional

需要修复该错误,优先实现无循环的向量化解法。

错误原因与循环修复方案

错误根源在于T[paths[:,s-1]]返回的是二维数组(形状为(n_sims, n_states)),但rng.choice的p参数仅接受对应单个样本的一维概率分布。以下是两种修复循环的可行方式:

方式1:逐样本调用choice

rng = np.random.default_rng(seed=123)
# 修正变量名错误:n_sim -> n_sims;循环范围改为1到n_steps(避免索引越界)
for s in range(1, n_steps):
    current_probs = T[paths[:, s-1]]
    # 给每个样本单独分配对应的概率分布进行采样
    paths[:, s] = np.array([rng.choice(states, p=p) for p in current_probs])

方式2:利用多项分布批量采样

rng = np.random.default_rng(seed=123)
for s in range(1, n_steps):
    current_probs = T[paths[:, s-1]]
    # 生成多项分布样本,取最大值索引即为选中状态
    paths[:, s] = np.argmax(rng.multinomial(1, current_probs, size=n_sims), axis=1)

向量化无循环实现

要彻底避免显式循环,可借助累积概率分布与随机数的广播比较来批量完成所有步骤的采样:

rng = np.random.default_rng(seed=123)

# 生成所有步骤所需的随机数:形状(n_sims, n_steps-1)
random_vals = rng.uniform(size=(n_sims, n_steps-1))

# 预处理转移矩阵的累积概率:形状(n_states, n_states)
cum_probs = T.cumsum(axis=1)

# 初始化路径数组
paths = np.zeros((n_sims, n_steps), dtype=int)

# 利用广播与argmax完成批量状态更新
for s in range(1, n_steps):
    # 获取当前所有样本对应的累积概率分布
    current_cum_probs = cum_probs[paths[:, s-1]]
    # 找到第一个大于随机数的索引,即为下一个状态
    paths[:, s] = np.argmax(current_cum_probs > random_vals[:, s-1, None], axis=1)

注:若追求极致无循环,可结合numba的scan函数实现,但上述代码在可读性与效率间已达成较好平衡,适合大多数场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 22:43:19