Keras get_weights循环调用耗时逐次递增的技术求助
解决Keras模型
get_weights()调用耗时随迭代增加的问题 听起来你遇到的这个问题挺典型的——在用遗传算法迭代生成Keras模型、直接用numpy处理权重融合/变异时,哪怕是相同规模的模型,get_weights()的调用耗时居然会越来越长。我之前做类似的神经架构搜索项目时也踩过这个坑,咱们一步步拆解可能的原因和解决办法:
可能的核心原因
- 模型实例与资源未被彻底回收:每一代生成100个模型后,如果只是简单丢弃引用,Python的垃圾回收机制可能没及时清理模型关联的计算图、张量缓存等资源,内存里堆积的残留状态会让后续
get_weights()的底层操作越来越耗时。 - numpy数组引用残留与内存碎片化:直接操作numpy权重数组时,若每次变异/融合都创建大量临时数组且未妥善释放,会导致内存碎片化,间接拖慢Keras读取权重的速度。
- Keras权重追踪机制的额外开销:即使你没在训练模型,Keras默认会追踪权重的更新状态、梯度信息等。每创建一个新模型,这些追踪状态就会累积,增加
get_weights()的调用成本。
针对性解决方案
1. 强制清理模型资源与垃圾回收
在每一代迭代结束后,主动清理模型并触发垃圾回收,这是最有效的第一步:
import gc from keras import backend as K # 处理完当前世代的所有模型后执行 for model in current_generation_models: # 清理Keras会话资源,比单纯del更彻底 K.clear_session() # 删除模型引用 del model # 强制触发Python垃圾回收 gc.collect()
K.clear_session()会释放当前TensorFlow/Keras会话的所有计算图和关联资源,避免残留状态堆积。
2. 优化权重处理的内存效率
避免创建不必要的临时numpy数组,尽量原地操作减少内存拷贝:
# 不推荐:每次都生成新数组 new_weights = parent_weights + mutation_strength * np.random.randn(*parent_weights.shape) # 推荐:原地修改数组,减少内存开销 parent_weights += mutation_strength * np.random.randn(*parent_weights.shape) # 用完临时变量后立即释放引用 del mutation_strength
另外,使用numpy的view()而非copy()(当不需要独立数组时),也能减少内存占用。
3. 关闭Keras不必要的追踪功能
创建模型时禁用不需要的动态追踪选项,减少状态维护开销:
from keras import Sequential, layers # 提前定义固定模型结构 base_model = Sequential([ layers.Dense(64, activation='relu'), layers.Dense(10, activation='softmax') ]) # 每代生成模型时,关闭eager执行(如果不需要动态图) new_model = base_model.clone_model() new_model.compile(run_eagerly=False) # 如果不需要训练,直接设置trainable=False减少状态追踪 new_model.trainable = False
clone_model()只复制模型结构,不复制权重和训练状态,能避免重复创建模型带来的额外开销。
4. 复用模型结构而非每次重建
与其每一代都从头创建全新模型,不如提前定义好固定结构,每次只重置权重:
# 提前定义一次结构 base_model = Sequential([...]) # 每代生成100个模型时复用结构 for _ in range(100): new_model = base_model.clone_model() # 直接设置随机生成的权重 new_model.set_weights(generate_random_weights(base_model)) # 处理模型...
这种方式能大幅减少模型初始化的资源消耗,从根源上降低get_weights()的潜在开销。
验证方法
你可以在每一代调用get_weights()前后,用psutil监控内存变化,确认是否是资源堆积问题:
import psutil print(f"调用前内存占用(MB):{psutil.Process().memory_info().rss / 1024 / 1024:.2f}") weights = model.get_weights() print(f"调用后内存占用(MB):{psutil.Process().memory_info().rss / 1024 / 1024:.2f}")
如果内存占用持续上升,说明资源回收不彻底,上面的清理方法就能解决问题。
内容的提问来源于stack exchange,提问作者Eric
相关产品推荐
相关产品推荐

