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

Keras get_weights/set_weights调用耗时过长原因及优化方案

问题描述

在实现迭代幅值剪枝实验时,每个训练迭代批次都会重复调用layer.get_weights()与layer.set_weights()方法。经测试,包含这两个调用的回调操作耗时达0.01s,而单批次训练本身仅耗时0.004s,直接导致总训练时长翻倍。
最初判断该操作仅涉及GPU侧的张量搬运,耗时不应与批次迭代过程中的大规模矩阵乘法处于同一量级,因此需要明确该耗时异常的成因,以及降低set_weights()、get_weights()调用耗时的可行方案。

相关代码

剪枝权重回调实现

### PRUNE WEIGHTS CALLBACK ###
class pruneModelCallback(Callback):
    def __init__(self, init_weight_dict=None, mask_dict=None):
        self.n_batches = 0
        self.init_weight_dict = init_weight_dict
        self.mask_dict = mask_dict
    
    def on_train_batch_begin(self, batch, logs=None):
        # 初始化时保存权重
        if self.n_batches == 0:
            if self.init_weight_dict is not None:
                for layer_i in range(len(self.model.layers)):
                    w = self.init_weight_dict['w_'+str(layer_i+1)]
                    b = self.init_weight_dict['b_'+str(layer_i+1)]

                    self.model.layers[layer_i].set_weights([w,b])
            else:
                self.init_weight_dict = {}
                for layer_i in range(len(self.model.layers)):
                    w = self.model.layers[layer_i].get_weights()[0]
                    b = self.model.layers[layer_i].get_weights()[1]

                    self.init_weight_dict['w_'+str(layer_i+1)] = w
                    self.init_weight_dict['b_'+str(layer_i+1)] = b
        
        self.n_batches = self.n_batches + 1
        
    # 存在性能问题的函数,每个训练批次都会执行
    def on_train_batch_end(self, batch, logs=None):
        # 将剪枝权重置零
        if self.mask_dict is not None:
            for layer_i in range(len(self.model.layers)):
                # 移除这两行调用可以小幅提升运行速度
                w = self.model.layers[layer_i].get_weights()[0]
                b = self.model.layers[layer_i].get_weights()[1]

                w_mask = self.mask_dict['w_'+str(layer_i+1)]

                # 掩码乘法本身耗时极低,移除该操作不会明显改变耗时
                w_pruned = w * w_mask

                # 移除该函数调用可以大幅提升运行速度
                self.model.layers[layer_i].set_weights([w_pruned,b])

运行时警告信息

1/629 [..............................] - ETA: 0s - loss: 2.3211 - accuracy: 0.0781 
WARNING:tensorflow:Callbacks method `on_train_batch_end` is slow compared to the batch time (batch time: 0.0040s vs `on_train_batch_end` time: 0.0100s). Check your callbacks.

实验场景说明

本次实验为复现彩票假说(Lottery Ticket Hypothesis)中提出的迭代幅值剪枝方法,实现逻辑为在每轮迭代中将剪枝掩码与权重做逐元素相乘,抵消权重更新的影响,保证被剪枝的权重始终为0,因此需要在每个训练迭代调用get_weights()和set_weights()。
实验所用模型为全部由全连接层构成的标准DNN,未使用批归一化、Dropout或其他正则化手段,模型定义代码如下:

model = Sequential([
    Dense(300, input_dim=input_dim[0], activation='relu'),
    Dense(100, activation='relu'),
    Dense(50, activation='relu'),
    Dense(output_dim[0], activation='softmax')
])

model.compile(
    optimizer = keras.optimizers.Adam(lr=1.2e-4),
    loss = tf.keras.losses.CategoricalCrossentropy(),
    metrics = ['accuracy']
)

迭代幅值剪枝流程本身需要多次重复训练模型,整体耗时较长,任何可行的加速方案都能显著提升实验效率。

耗时异常核心成因
  • get_weights()和set_weights()并非纯GPU侧操作:两个方法调用时会强制触发CPU-GPU同步,get_weights()会将GPU显存中的张量拷贝到主机内存并转换为numpy数组,set_weights()会将主机内存的numpy数组拷贝回GPU显存,每次调用都会插入同步点,直接打断GPU的异步计算流水线,同步本身的开销远大于张量数据拷贝的开销。
  • 高频调用放大同步开销:GPU训练默认采用异步流水线执行逻辑,频繁的同步点会导致GPU频繁等待CPU调度,无法将多个批次的计算排队重叠执行,在单批次训练耗时较短的场景下,同步开销的占比会被极度放大,这也是当前场景下回调耗时超过训练本身的核心原因。
  • 附加CPU侧校验开销:每次调用get_weights()/set_weights()时,Keras内部会串行执行权重形状校验、数值合法性检查、训练状态追踪等CPU侧逻辑,高频调用下这些零散操作的累积耗时也不可忽略。
可行优化方案
  • 方案1:全GPU侧实现掩码逻辑,彻底规避主机拷贝
    这是性能最优的方案,完全抛弃回调中手动读写numpy权重的实现方式:可以将剪枝掩码设置为模型中不可训练的tf.Variable,通过自定义权重约束(Constraint)的方式,让优化器每次更新权重后,自动在GPU侧完成权重和掩码的逐元素相乘,全程不触发GPU-CPU同步,也不需要手动调用get_weights()/set_weights(),开销可以降低两个数量级以上。
  • 方案2:降低掩码应用频率
    迭代幅值剪枝不需要每个批次都强制应用掩码,被剪枝的权重在少量批次中产生的微小更新对最终剪枝结果的影响可以忽略。可以将掩码应用逻辑从on_train_batch_end移到on_epoch_end,或者每N个批次应用一次掩码,调用开销会随调用频率降低同比下降。
  • 方案3:直接操作TensorFlow变量,跳过numpy转换环节
    如果需要保留批次级更新的逻辑,不要通过get_weights()获取numpy格式的权重,直接引用层自带的layer.kernel、layer.bias张量对象,用TensorFlow算子在GPU侧完成掩码乘法,再通过assign()方法直接更新变量值,全程不触发同步和numpy转换。
    原有逻辑的替换示例如下:
    # 原实现(触发CPU-GPU同步,开销高)
    # w = layer.get_weights()[0]
    # w_pruned = w * w_mask
    # layer.set_weights([w_pruned, b])
    
    # 新实现(全GPU执行,无同步开销)
    w_mask = tf.constant(self.mask_dict['w_'+str(layer_i+1)], dtype=layer.kernel.dtype)
    layer.kernel.assign(layer.kernel * w_mask)
    

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 05:57:17