使用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
相关产品推荐
相关产品推荐

