TensorFlow微调BERT用tf.GradientTape自定义训练循环报错如何解决?
问题原因
该报错是TensorFlow 2.3版本的已知缺陷,核心原因是从TF Hub加载的可训练KerasLayer在eager执行模式的自定义训练循环中,部分内部参数的资源句柄没有被梯度带正确识别。而model.fit默认会将训练逻辑包装为图执行模式,会主动完成所有参数的初始化和句柄注册,所以不会触发该报错。
解决方案
- 方案1(最便捷,推荐):给自定义训练步添加
@tf.function装饰器,对齐model.fit的执行逻辑,触发图模式下的参数自动注册,修改后的训练步代码如下:
@tf.function def train_step(train_batch): train_x, train_y = train_batch with tf.GradientTape() as tape: ypred = model(train_x) loss = tf.reduce_mean(tf.keras.losses.binary_crossentropy(train_y, ypred)) grads = tape.gradient(loss, model.trainable_weights) optimizer.apply_gradients(zip(grads, model.trainable_weights)) return loss
- 方案2(需要eager模式调试时使用):在定义BERT层之后手动调用
build方法强制初始化所有参数句柄,修改模型定义中BERT层相关代码:
bert_layer = hub.KerasLayer("https://tfhub.dev/tensorflow/bert_en_uncased_L-12_H-768_A-12/2", trainable=True) # 新增如下行,传入三个输入的对应shape bert_layer.build([(None, 512), (None, 512), (None, 512)]) pooled_output, sequence_output = bert_layer([input_word_ids, input_mask, input_type_ids])
- 方案3(长期兼容):将TensorFlow版本升级到2.4.0及以上,该版本已经修复了TF Hub层在eager模式下的参数跟踪缺陷。
内容的提问来源于stack exchange,提问作者Jane Sully
相关产品推荐
相关产品推荐

