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

如何修复R语言中Q-Learning算法的sample函数概率参数错误?

修复R语言Q-Learning实现中的概率参数错误

原始代码与错误

用户实现Q-Learning的代码如下:

# Define the map
map <- matrix(c(0, 1, 1, 0, 0, 0, 0, 1), nrow = 2, ncol = 4, byrow = TRUE)
# State labels
rownames(map) <- c("Start", "End")
# Action labels
colnames(map) <- c("Up", "Down", "Left", "Right")
# Rewards for each state-action pair
rewards <- matrix(c(-1, -1, -1, -1, -1, -1, -1, 10), nrow = 2, ncol = 4, byrow = TRUE)

# Q-Learning Algorithm
q_learning <- function(P, R, gamma = 0.9, alpha = 0.1, epsilon = 0.1, max_iter = 1000) {
  # Initialize the Q-value function
  Q <- matrix(rep(0, nrow(P) * ncol(P)), nrow = nrow(P), ncol = ncol(P))
  # Initialize the state
  state <- sample(1:nrow(P), 1)
  # Iterate until convergence or maximum iterations reached
  for (i in 1:max_iter) {
    # Choose an action using epsilon-greedy policy
    if (runif(1) < epsilon) {
      action <- sample(1:ncol(P), 1)
    } else {
      action <- which.max(Q[state, ])
    }
    # Observe the next state and reward
    prob <- P[state, action]
    next_state <- sample(1:nrow(P), 1, prob = prob)
    reward <- R[state, action]
    # Update the Q-value function
    Q[state, action] <- Q[state, action] + alpha * (reward + gamma * max(Q[next_state, ]) - Q[state, action])
    # Update the state
    state <- next_state
  }
  # Derive the optimal policy (argmax in R using the which.max)
  policy <- apply(Q, 1, which.max)
  # Return the Q-value function and policy
  return(list(Q = Q, policy = policy))
}


# Run the Q-Learning Algorithm on the map
q_learning(P = map, R = rewards, gamma = 0.9, alpha = 0.1, epsilon = 0.1, max_iter = 1000)

运行后触发错误:

Error in sample.int(length(x), size, replace, prob) :
incorrect number of probabilities

错误原因

问题出在转移概率矩阵的定义和使用逻辑不匹配:

  • 你把map当作转移概率矩阵P传入函数,但P[state, action]取出的是单个数值,而sample函数的prob参数需要一个长度与采样对象数量一致的概率向量(这里要从2个状态中采样,所以需要长度为2的向量)。
  • 原map矩阵的结构错误,它没有正确表示“当前状态-动作”对应的所有下一状态的概率分布。

修复方案

1. 重新定义转移概率矩阵

将P定义为三维数组,维度为[状态数, 动作数, 下一状态数],用来存储每个状态-动作对下,转移到各个下一状态的概率。针对2状态4动作场景:

# 定义转移概率数组:[状态, 动作, 下一状态]
P <- array(0, dim = c(2, 4, 2))
# Start状态(索引1)的动作转移规则
P[1, , ] <- matrix(c(
  0, 1,  # Up动作:100%转移到End状态(索引2)
  1, 0,  # Down动作:100%留在Start状态(索引1)
  1, 0,  # Left动作:100%留在Start状态(索引1)
  0, 1   # Right动作:100%转移到End状态(索引2)
), nrow=4, byrow=TRUE)
# End状态(索引2)的动作转移规则:所有动作都留在End状态
P[2, , ] <- matrix(c(1,0, 1,0, 1,0, 1,0), nrow=4, byrow=TRUE)

# 添加状态与动作标签
dimnames(P) <- list(
  state = c("Start", "End"),
  action = c("Up", "Down", "Left", "Right"),
  next_state = c("Start", "End")
)

# 奖励矩阵保持不变
rewards <- matrix(c(-1, -1, -1, -1, -1, -1, -1, 10), nrow = 2, ncol = 4, byrow = TRUE)
rownames(rewards) <- c("Start", "End")
colnames(rewards) <- c("Up", "Down", "Left", "Right")

2. 修改Q-Learning函数的转移概率获取逻辑

调整函数中获取prob的代码,从三维数组中取出对应状态-动作的下一状态概率向量:

q_learning <- function(P, R, gamma = 0.9, alpha = 0.1, epsilon = 0.1, max_iter = 1000) {
  n_states <- dim(P)[1]
  n_actions <- dim(P)[2]
  
  # Initialize the Q-value function
  Q <- matrix(rep(0, n_states * n_actions), nrow = n_states, ncol = n_actions)
  # Initialize the state
  state <- sample(1:n_states, 1)
  # Iterate until convergence or maximum iterations reached
  for (i in 1:max_iter) {
    # Choose an action using epsilon-greedy policy
    if (runif(1) < epsilon) {
      action <- sample(1:n_actions, 1)
    } else {
      action <- which.max(Q[state, ])
    }
    # Observe the next state and reward
    prob <- P[state, action, ]  # 获取当前状态-动作对应的所有下一状态概率向量
    next_state <- sample(1:n_states, 1, prob = prob)
    reward <- R[state, action]
    # Update the Q-value function
    Q[state, action] <- Q[state, action] + alpha * (reward + gamma * max(Q[next_state, ]) - Q[state, action])
    # Update the state
    state <- next_state
  }
  # Derive the optimal policy
  policy <- apply(Q, 1, which.max)
  # 给策略添加状态标签
  names(policy) <- rownames(R)
  # Return the Q-value function and policy
  return(list(Q = Q, policy = policy))
}

3. 运行修复后的代码

# 运行算法
result <- q_learning(P = P, R = rewards, gamma = 0.9, alpha = 0.1, epsilon = 0.1, max_iter = 1000)

# 查看结果
print(result$Q)
print(result$policy)

这样就能正确运行,不会再触发概率数量不匹配的错误。


内容的提问来源于stack exchange,提问作者Homer Jay Simpson

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 19:07:41