在线Categorical Temporal-Difference学习:概率和不为1问题排查
在线Categorical Temporal-Difference算法实现问题排查
问题描述
正在阅读《MIT分布强化学习》一书,实现了用于学习策略回报分布的在线Categorical Temporal-Difference学习算法。运行测试代码后发现最终回报的概率和不为1,且预期概率应集中在locations = np.linspace(-1, 1, M)的正值区域,但实际不符合预期。
实现代码如下:
import numpy as np import gymnasium as gym def online_CTD(policy): EPISODE = 10000 LEARNING_RATE = 0.05 DISCOUNT = 0.99 EXPLORE = 0.1 MAX_STEPS = 1000 M = 31 env = gym.make('Taxi-v3') # initialise locations = np.linspace(-1, 1, M) #policy = np.random.rand(env.observation_space.n, env.action_space.n) probabilities = np.ones((M, env.observation_space.n)) probabilities = probabilities / np.sum(probabilities) for e in range(EPISODE): terminated = False truncated = False state, _ = env.reset() while not terminated and not truncated: #space = policy[state] #action = np.random.choice(len(space), p=space) action = np.argmax(policy[state]) next_state, reward, terminated, truncated, _ = env.step(action) temp_policy = np.array([0 for i in range(M)]) for j in range(1, M+1): if terminated or truncated: g = reward else: g = reward + DISCOUNT * locations[j - 1] if g <= locations[0]: temp_policy[0] += probabilities[j - 1, next_state] elif g >= locations[-1]: temp_policy[-1] += probabilities[j - 1, next_state] else: i = np.searchsorted(locations, g) - 1 # print(locations, g, i) d = (g - locations[i]) / (locations[i+1] - locations[i]) temp_policy[i] += (1-d) * probabilities[j - 1, next_state] temp_policy[i+1] += d * probabilities[j - 1, next_state] for i in range(M): probabilities[i, state] = (1-LEARNING_RATE) * probabilities[i, state] + LEARNING_RATE * temp_policy[i] state = next_state return probabilities
测试代码:
probabilities = online_CTD(Q_table) dist = np.sum(probabilities, axis=1) print(dist) print(np.sum(dist))
问题分析与修正方案
1. 概率初始化错误
当前代码将整个二维概率数组的总和归一化,不符合Categorical TD的核心要求——每个状态对应的M个分位点概率之和必须为1。
- 错误代码:
probabilities = np.ones((M, env.observation_space.n)) probabilities = probabilities / np.sum(probabilities) - 修正代码(按列归一化,每个状态的概率初始均匀分布):
probabilities = np.ones((M, env.observation_space.n)) / M
2. 缩进错误导致中间区间概率更新失效
在处理g处于locations中间区间的逻辑时,d的计算和temp_policy的更新代码缩进错误,导致这部分逻辑不在else块内,既会引发未定义变量错误,又会跳过中间区间的概率分配逻辑。
- 错误缩进部分:
else: i = np.searchsorted(locations, g) - 1 # print(locations, g, i) d = (g - locations[i]) / (locations[i+1] - locations[i]) temp_policy[i] += (1-d) * probabilities[j - 1, next_state] temp_policy[i+1] += d * probabilities[j - 1, next_state] - 修正后(将
d及后续代码缩进至else块内):else: i = np.searchsorted(locations, g) - 1 # print(locations, g, i) d = (g - locations[i]) / (locations[i+1] - locations[i]) temp_policy[i] += (1-d) * probabilities[j - 1, next_state] temp_policy[i+1] += d * probabilities[j - 1, next_state]
3. 分位点范围与环境回报不匹配
Taxi-v3环境的实际回报范围远超出[-1,1]:每步会有-1的惩罚,成功时获得+20回报,若触发截断(MAX_STEPS=1000)总回报可达-1000。当前设置的locations范围过小,导致大部分回报值被强制映射到边界分位点,且无法覆盖正回报的合理区间。
- 修正方案:根据环境实际回报范围调整
locations,例如:locations = np.linspace(-1000, 20, M)
4. 可选:添加概率归一化步骤(缓解数值误差)
在每次更新probabilities后,可对每个状态的概率进行归一化,避免长期学习过程中因数值累积误差导致概率和偏离1:
# 在更新probabilities后添加 probabilities[:, state] /= np.sum(probabilities[:, state])
内容的提问来源于stack exchange,提问作者YUTONG JI
相关产品推荐
相关产品推荐

