Chuyển đến nội dung chính

第 2 課:動態規劃-策略迭代與值迭代

貝爾曼方程式。政策評估、政策改進。值迭代算法。 Python 中的 GridWorld 實作。

🧠 人工智慧與機器學習 — 第 1 課 第 2 課:動態規劃 — 策略 迭代與值迭代

強化學習:從基礎到高級

第 1 部分:強化學習基礎 — 馬可夫決策過程與表格方法

亞洲開發網

簡介

動態規劃 (DP) 在了解整個模型的情況下求解 MDP — 轉移機率 P(s'|s,a) 和獎勵 R。這是所有 RL 演算法的理論基礎。


1. 貝爾曼方程

貝爾曼期望方程

$$V^\pi(s) = \sum_a \pi(a|s) \sum_{s'} P(s'|s,a)[R(s,a,s') + \gamma V^\pi(s')]$$

貝爾曼最優方程

$$V^(s) = \max_a \sum_{s'} P(s'|s,a)[R(s,a,s') + \gamma V^(s')]$$


2.政策評估

計算固定策略的 V(s):

import numpy as np

def policy_evaluation(policy, env, gamma=0.99, theta=1e-8):
    V = np.zeros(env.nS)
    while True:
        delta = 0
        for s in range(env.nS):
            v = 0
            for a, action_prob in enumerate(policy[s]):
                for prob, next_state, reward, done in env.P[s][a]:
                    v += action_prob * prob * (reward + gamma * V[next_state])
            delta = max(delta, abs(V[s] - v))
            V[s] = v
        if delta < theta:
            break
    return V

3. 政策改進

基於V(s)的貪婪改進:

def policy_improvement(V, env, gamma=0.99):
    policy = np.zeros([env.nS, env.nA])
    for s in range(env.nS):
        q_values = np.zeros(env.nA)
        for a in range(env.nA):
            for prob, next_state, reward, done in env.P[s][a]:
                q_values[a] += prob * (reward + gamma * V[next_state])
        best_action = np.argmax(q_values)
        policy[s][best_action] = 1.0
    return policy

4. 策略迭代

def policy_iteration(env, gamma=0.99):
    policy = np.ones([env.nS, env.nA]) / env.nA  # Uniform random
    while True:
        V = policy_evaluation(policy, env, gamma)
        new_policy = policy_improvement(V, env, gamma)
        if np.array_equal(policy, new_policy):
            break
        policy = new_policy
    return policy, V

5. 值迭代

def value_iteration(env, gamma=0.99, theta=1e-8):
    V = np.zeros(env.nS)
    while True:
        delta = 0
        for s in range(env.nS):
            v = V[s]
            V[s] = max(
                sum(p * (r + gamma * V[s_])
                    for p, s_, r, _ in env.P[s][a])
                for a in range(env.nA)
            )
            delta = max(delta, abs(v - V[s]))
        if delta < theta:
            break
    return V

6.GridWorld 演示

import gymnasium as gym

env = gym.make("FrozenLake-v1", is_slippery=False)
policy, V = policy_iteration(env.unwrapped, gamma=0.99)

print("Optimal Value Function:")
print(V.reshape(4, 4))
print("Optimal Policy (0=L, 1=D, 2=R, 3=U):")
print(np.argmax(policy, axis=1).reshape(4, 4))

總結

方法方法收斂複雜度
政策迭代評估→改進幾次迭代每次評估的 O(S²A)
價值迭代一步前瞻性多次迭代每次迭代 O(SA)