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

基于TensorFlow的SARSA算法实现问题咨询(CartPole案例)

排查SARSA算法在CartPole实现中的学习效果问题

我看了你的SARSA实现代码,确实存在几个关键问题导致学习效果受限,咱们一步步拆解并给出修正方案:

核心问题分析

1. 单样本训练带来的高方差

你当前每一步都用单个状态-动作对计算梯度并更新网络,这种方式的方差极大,神经网络很难稳定学习到有效的Q值——就像每次只看一个样本就调整模型,很容易在噪声里来回震荡。

2. 生硬的ε衰减策略

只在第100个episode把ε从0.3直接降到0.1,之后固定不变的做法太突兀,不利于模型平滑从"探索"过渡到"利用"阶段,容易错过更优的学习节奏。

3. 目标Q值的不稳定波动

非终止状态下,你直接用当前训练的网络计算目标Q值,这会导致目标值随着网络更新不断变化,进一步加剧训练的不稳定性。

4. 小细节:奖励累计的通用性问题

你的sum_returns += 1在CartPole环境中没问题(每步奖励都是1),但改成sum_returns += reward能适配更多奖励非固定的环境。

修正后的完整代码

import numpy as np
import matplotlib.pylab as plt
import random
import gym
import tensorflow as tf
from collections import deque

# 经验回放缓冲区:存储轨迹数据,批量采样降低训练方差
class ReplayBuffer:
    def __init__(self, capacity):
        self.buffer = deque(maxlen=capacity)
    
    def store(self, state, action, reward, next_state, next_action, done):
        self.buffer.append((state, action, reward, next_state, next_action, done))
    
    def sample(self, batch_size):
        batch = random.sample(self.buffer, batch_size)
        states, actions, rewards, next_states, next_actions, dones = zip(*batch)
        return np.array(states), np.array(actions), np.array(rewards), np.array(next_states), np.array(next_actions), np.array(dones)
    
    def __len__(self):
        return len(self.buffer)

# 构建Q网络:增加一层隐藏层提升拟合能力
def build_network():
    return tf.keras.Sequential([
        tf.keras.layers.Dense(16, activation='relu', input_shape=[4]),
        tf.keras.layers.Dense(16, activation='relu'),
        tf.keras.layers.Dense(2)
    ])

# 主网络(负责更新)和目标网络(负责计算稳定的目标Q值)
main_net = build_network()
target_net = build_network()
target_net.set_weights(main_net.get_weights())  # 初始化目标网络参数与主网络一致

def q_value(network, states, actions):
    q_vals = network(tf.convert_to_tensor(states, dtype=tf.float32))
    return tf.gather(q_vals, actions, axis=1)  # 批量提取对应动作的Q值

def policy(network, state, epsilon):
    if np.random.rand() < epsilon:
        return random.choice([0, 1])
    else:
        q_vals = network(tf.convert_to_tensor([state], dtype=tf.float32))
        return np.argmax(q_vals.numpy())

# 初始化环境和缓冲区
env = gym.make('CartPole-v0')
buffer = ReplayBuffer(capacity=10000)

# 调整后的超参数
discount = 0.99
optimizer = tf.keras.optimizers.Adam(learning_rate=0.001)
episodes = 1000
epsilon_start = 0.3
epsilon_end = 0.01
epsilon_decay = 0.995  # 指数衰减,平滑降低探索概率
batch_size = 32
target_update_freq = 50  # 每50个episode更新一次目标网络

rewards = []

for episode in range(episodes):
    state = env.reset()
    sum_returns = 0
    # 平滑衰减ε
    epsilon = max(epsilon_end, epsilon_start * (epsilon_decay ** episode))
    
    while True:
        action = policy(main_net, state, epsilon)
        next_state, reward, done, _ = env.step(action)
        next_action = policy(main_net, next_state, epsilon)
        sum_returns += reward
        
        # 存储轨迹到经验回放缓冲区
        buffer.store(state, action, reward, next_state, next_action, done)
        
        # 缓冲区有足够数据时才开始批量训练
        if len(buffer) >= batch_size:
            states, actions, rewards_batch, next_states, next_actions, dones = buffer.sample(batch_size)
            
            with tf.GradientTape() as tape:
                tape.watch(main_net.trainable_variables)
                current_q = q_value(main_net, states, actions)
                # 用目标网络计算稳定的目标Q值
                next_q = q_value(target_net, next_states, next_actions)
                target_q = rewards_batch + discount * next_q * (1 - dones)
                loss = tf.reduce_mean(tf.square(target_q - current_q))
            
            gradients = tape.gradient(loss, main_net.trainable_variables)
            optimizer.apply_gradients(zip(gradients, main_net.trainable_variables))
        
        # 定期更新目标网络参数
        if episode % target_update_freq == 0:
            target_net.set_weights(main_net.get_weights())
        
        if done:
            break
        
        state = next_state
    
    rewards.append(sum_returns)
    # 每20个episode打印进度
    if episode % 20 == 0:
        print(f"Episode {episode}, Average Reward (last 20): {np.mean(rewards[-20:]):.2f}, Epsilon: {epsilon:.3f}")

# 绘制奖励曲线
plt.plot(np.arange(episodes), rewards)
plt.title('SARSA with Experience Replay and Target Network')
plt.xlabel('Episode')
plt.ylabel('Total Reward')
plt.show()

关键改进说明

  1. 经验回放:通过批量采样历史轨迹数据训练,大幅降低梯度方差,让模型学习更稳定。
  2. 平滑ε衰减:用指数衰减策略让探索概率从0.3逐步降到0.01,平衡探索与利用的节奏。
  3. 目标网络:分离主网络和目标网络,用目标网络计算稳定的目标Q值,避免目标值随网络更新频繁波动。
  4. 增强网络结构:增加一层隐藏层,提升模型对复杂状态的拟合能力。

运行修正后的代码,你应该能看到奖励值逐渐上升并稳定在CartPole的最大奖励(200)附近。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.11 07:50:09