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

DQN训练中keras Model.fit调用导致内存持续增长问题

DQN训练时model.fit()内存泄漏问题

问题描述

运行深度Q网络(DQN)时,每次调用model.fit()都会导致程序内存占用持续增加,长时间运行后内存会被完全耗尽。当前环境为TensorFlow 2.12.0、Python 3.10,运行在CPU平台。已尝试将训练数据x、y转换为张量,调用垃圾回收(gc)以及keras.backend.clear_session(),但均未解决内存泄漏问题。

memory_profiler对train()函数的内存分析结果

Line #    Mem usage    Increment  Occurrences   Line Contents
=============================================================
   124    315.7 MiB    315.7 MiB           1       @profile
   125                                             def train(self):
   126                                         
   127    315.7 MiB      0.0 MiB           1           if len(self.replaybuffer) < MIN_REPLAY_BUFFER:
   128                                                     return
   129                                         
   130                                                 #get idexes from replay buffer
   131    315.7 MiB      0.0 MiB           1           rpbIdexes = np.random.choice(range(len(self.replaybuffer)), size=BATCHSIZE, replace=False)

   132                                         
   133                                                 #predict in bulk for speed up
   134    315.7 MiB      0.0 MiB          35           nextObservations = np.array([self.replaybuffer[rpbId][4] for rpbId in rpbIdexes])

   135    315.7 MiB      0.0 MiB           1           if self.useTarget:
   136    315.7 MiB      0.0 MiB           1               nextPredictions = self.targetmodel(nextObservations.reshape(-1, 4))

   137                                                 else:
   138                                                     nextPredictions = self.model(nextObservations.reshape(-1, 4))

   139                                         
   140                                                 # x and y arrays for fitting
   141    315.7 MiB      0.0 MiB           1           x, y = [], []
   142    315.7 MiB      0.0 MiB          33           for idx, rpbId in enumerate(rpbIdexes):
   143    315.7 MiB      0.0 MiB          32               observation = self.replaybuffer[rpbId][0]
   144    315.7 MiB      0.0 MiB          32               x.append(observation)
   145                                         
   146    315.7 MiB      0.0 MiB          32               nextPrediction = np.max(nextPredictions[idx])

   147                                         
   148    315.7 MiB      0.0 MiB          32               done = self.replaybuffer[rpbId][5]
   149    315.7 MiB      0.0 MiB          32               reward = self.replaybuffer[rpbId][3]
   150    315.7 MiB      0.0 MiB          32               if done:
   151    315.7 MiB      0.0 MiB           2                   target = reward
   152                                                     else:
   153    315.7 MiB      0.0 MiB          30                   target = reward + self.gamma * nextPrediction

   154                                                     # change the prediciton
   155    315.7 MiB      0.0 MiB          32               takenAction = self.replaybuffer[rpbId][2]
   156    315.7 MiB      0.0 MiB          32               prediction = self.replaybuffer[rpbId][1]
   157                                         
   158    315.7 MiB      0.0 MiB          32               prediction = prediction.numpy()
   159    315.7 MiB      0.0 MiB          32               prediction[takenAction] = target
   160    315.7 MiB      0.0 MiB          32               y.append(prediction)
   161    315.7 MiB      0.0 MiB           1           x = tf.convert_to_tensor(x)
   162    315.7 MiB      0.0 MiB           1           y= tf.convert_to_tensor(y)
   163    316.1 MiB      0.4 MiB           1           history = self.model.fit(x=x, y=y, batch_size=BATCHSIZE, verbose=0, shuffle=False)

   164    316.1 MiB      0.0 MiB           1           return history

可行的解决思路

  • 用tf.function装饰train函数:将训练逻辑转为图模式执行,减少即时执行模式下反复创建计算图节点带来的内存残留。
  • 复用训练数据容器:不要每次训练都新建x、y列表,提前初始化固定大小的numpy数组或张量,直接修改数据内容,降低内存碎片化。
  • 清理训练历史对象:model.fit()返回的history会存储训练指标,每次训练后可以手动清空history的属性,或者不保留该对象的引用。
  • 优化 replay buffer 存储:将replaybuffer中存储的Tensor类型prediction提前转为numpy数组,避免每次训练时重复调用numpy()产生不必要的内存开销。
  • 调整TensorFlow版本:TensorFlow 2.12.x存在已知的内存泄漏问题,可尝试降级到2.11.x或升级到2.15.x等稳定版本。
  • 显式触发垃圾回收:在model.fit()执行完成后,立刻调用gc.collect(),强制回收未被引用的内存对象。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.25 18:12:45