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

使用pyGad训练神经网络时内存持续上涨崩溃问题咨询

解决pyGad训练神经网络时的内存持续增长问题

针对你遇到的每代训练后内存占用攀升、最终崩溃的问题,可以通过以下几个方向优化:

1. 手动清理TensorFlow计算图与显存

TensorFlow在每次预测后可能残留未释放的计算图节点或张量,导致内存累积。在fitness_func的最后添加显存清理代码:

import gc
import tensorflow as tf
import datetime

def fitness_func(ga_instance, solution, solution_idx):
    global train, label, model
    t0 = datetime.datetime.now()
    preds = pygad.kerasga.predict(model=model, solution=solution, data=train, verbose=0, batch_size=2**13)
    t1 = datetime.datetime.now()
    print(t1 - t0)
    
    scores = label[preds>0.7].mean() - label[preds<0.3].mean()
    score = scores.mean()
    
    # 清理显存与计算图,重新构建模型
    tf.keras.backend.clear_session()
    input_layer = tf.keras.layers.Input(shape=(train.shape[1]))
    dense_layer1 = tf.keras.layers.Dense(units=train.shape[1], activation=tf.keras.layers.LeakyReLU(alpha=0.01))(input_layer)
    output_layer = tf.keras.layers.Dense(units=1)(dense_layer1)
    model = tf.keras.Model(inputs=input_layer, outputs=output_layer)
    
    # 删除临时变量并触发垃圾回收
    del preds, scores
    gc.collect()
    
    return score

2. 避免全局变量依赖

全局变量train、model可能导致内存无法被垃圾回收器正确释放,建议转为局部引用并在使用后主动清理:

def fitness_func(ga_instance, solution, solution_idx):
    # 局部引用全局变量
    local_train = train
    local_label = label
    local_model = model
    
    # ... 预测和计算得分代码 ...
    
    # 解除局部引用并回收内存
    del local_train, local_label, local_model, preds, scores
    gc.collect()
    return score

3. 手动分批预测替代内置predict

pyGad的kerasga.predict可能存在内存泄漏问题,改为手动分批处理数据,每批处理后立即释放内存:

def fitness_func(ga_instance, solution, solution_idx):
    global train, label, model
    batch_size = 2**13
    preds = []
    
    # 手动分批处理
    for i in range(0, len(train), batch_size):
        batch = train[i:i+batch_size]
        batch_pred = pygad.kerasga.predict(model=model, solution=solution, data=batch, verbose=0)
        preds.extend(batch_pred)
        # 清理当前批次的临时变量
        del batch, batch_pred
        gc.collect()
    
    preds = np.array(preds)
    scores = label[preds>0.7].mean() - label[preds<0.3].mean()
    score = scores.mean()
    
    del preds, scores
    gc.collect()
    return score

4. 调整pyGad的内存相关参数

  • 减少num_solutions:当前设置为5,可尝试降至3,减少每代需要存储的解数量
  • 关闭不必要的历史缓存:pyGad默认保留每代最优解,可通过参数关闭或减少保留数量
ga_instance = pygad.GA(num_generations=50, 
                       num_parents_mating=2, 
                       initial_population=keras_ga.population_weights, 
                       fitness_func=fitness_func,
                       on_generation=on_generation,
                       save_best_solutions=False,  # 关闭历史最优存储
                       keep_elitism=1)  # 仅保留1个精英解

5. 每代结束后强制垃圾回收

在on_generation函数末尾添加垃圾回收触发,确保每代结束后清理内存:

def on_generation(ga_instance):
    print(ga_instance.generations_completed, f"Fitness    = {ga_instance.best_solution()[1]}")
    gc.collect()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 21:52:34