如何手动设置Keras网络参数,实现比model.set_weights()更高效的方案?
优化方案
你遇到的性能问题核心来自两方面:一是set_weights本身冗余校验多,且每次调用都要完成CPU到GPU的数据拷贝;二是单染色体循环评估的模式完全没有利用GPU的并行计算能力。可通过以下方案逐步优化:
- 直接用
assign替换set_weights,减少冗余开销set_weights内部会做参数形状、格式、设备匹配的大量校验,本身开销很高,你可以直接遍历模型的可训练权重变量,直接赋值跳过校验,性能可以提升30%~50%:
def set_params(model: keras.Model, chromosome: np.ndarray, dummy_params): i = 0 for weight_var, layer_template in zip(model.trainable_weights, dummy_params): layer_size = layer_template.size reshaped_weight = np.reshape(chromosome[i:i+layer_size], layer_template.shape) weight_var.assign(reshaped_weight) i += layer_size
- 全流程静态图编译,消除CPU/GPU数据来回拷贝
你当前代码中predicted_outputs.numpy()、numpy侧计算损失的操作会导致数据在CPU和GPU之间来回传输,同时eager模式单步执行的开销很高。可以做如下修改:- 提前把输入、目标输出都转为Tensor存储在GPU上
- 用
tf.function包装整个评估逻辑,编译成静态图执行 - 损失计算完全在Tensor侧完成,最后统一转numpy输出
以上修改可以把前向传播和损失计算的性能提升2~3倍。
- 批量并行评估整个种群,消除循环开销
如果你的种群规模较大,逐个评估的模式浪费了GPU的并行能力,收益最高的优化是把整个种群的评估逻辑向量化:把模型权重扩展出种群大小的维度,输入也对应扩展,一次性完成所有染色体的前向传播和损失计算,完全消除循环调用set_weights的开销,整体评估速度可以提升5~10倍。 - 小优化:缓存权重形状信息
你当前代码中提前获取的dummy_params不需要保留权重值,只需要缓存各层权重的形状和大小即可,减少不必要的参数拷贝开销。
内容的提问来源于stack exchange,提问作者Emil Jansson
相关产品推荐
相关产品推荐

