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

基于TensorFlow的Q-learning模型训练速度下降问题求助

问题:TensorFlow实现Q-learning训练速度逐轮下降

我用TensorFlow实现带CNN架构的Q-learning智能体时,遇到训练速度逐轮显著降低的问题。每轮训练结束后会保存模型,但后续训练直接基于当前模型继续,不加载已保存的模型。精简后的核心代码如下:

import numpy as np
import tensorflow as tf

MAX_EPISODES = 50
CONTINUE = True

class QLearningAgent:
    def __init__(self, state_size, action_size):
        self.state_size = state_size
        self.action_size = action_size
        self.epsilon = 0.9
        self.epsilon_decay = 0.995
        self.epsilon_min = 0.1
        self.learning_rate = 0.01
        self.gamma = 0.95
        self.model = self.build_model()

    def build_model(self):
        # 自定义CNN模型架构
        # ...

    def act(self, state):
        # ε-贪心策略选择动作
        # ...

    def train(self, state, action, reward, next_state, done):
        # Q-learning训练逻辑
        # ...

    def save_model(self, filename):
        self.model.save(filename)

    def update_epsilon(self):
        self.epsilon = max(self.epsilon * self.epsilon_decay, self.epsilon_min)


env = BallEnvironment(max_steps=1000)

for episode in range(MAX_EPISODES):
    obs = env.reset()

    while True:
        env.render()

        left_action = env.left_ball.q_agent.act(np.reshape(obs, [1, *env.state_size]))
    
        next_obs, rewards, done, _ = env.step(left_action, right_action)

        left_state = np.reshape(obs, [1, *env.state_size])
        left_next_state = np.reshape(next_obs, [1, *env.state_size])
        env.left_ball.q_agent.train(left_state, left_action, rewards[0], left_next_state, done)

        obs = next_obs

        if done:
            env.left_ball.q_agent.save_model("left_trained_agent.h5")
            break

env.close()
解决方案建议
  • 降低模型保存开销:model.save()会完整序列化模型结构、权重和优化器状态,每轮执行会累积大量IO和序列化开销。可以改成:

    • 定期保存(比如每10轮保存一次),减少保存频率
    • 只保存权重(model.save_weights()),相比完整模型保存开销大幅降低
    • 如果不需要恢复优化器状态,保存时指定include_optimizer=False
  • 固化计算图:如果train方法中的逻辑是动态构建的,会导致TensorFlow计算图不断膨胀,拖慢后续训练。用@tf.function装饰train方法,固化计算图,避免重复构建:

    @tf.function
    def train(self, state, action, reward, next_state, done):
        # Q-learning训练逻辑
        # ...
    
  • 关闭训练时的渲染:env.render()在每步执行会占用大量CPU/GPU资源,训练阶段可以关闭渲染,只在需要验证或观察时开启:

    # 训练时注释或删除这行
    # env.render()
    
  • 清理TensorFlow会话:每轮训练结束后,调用tf.keras.backend.clear_session()清理未使用的计算节点和张量,避免内存泄漏:

    if done:
        env.left_ball.q_agent.save_model("left_trained_agent.h5")
        tf.keras.backend.clear_session()
        break
    
  • 检查优化器状态:如果使用Adam等自适应优化器,其内部状态(如动量项)会随训练累积,虽然这是正常行为,但如果保存模型时附带优化器状态,会增加保存开销。如果不需要恢复训练进度时的优化器状态,保存时排除优化器即可。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 03:12:19