TensorFlow中对应PyTorch model.eval()+no_grad()的等效实现
TensorFlow 复现BERT嵌入提取的对应实现
核心API对应关系
model.eval()等效:TensorFlow 没有全局切换模型推理/训练模式的独立方法,在前向传播调用模型时传入training=False即可,会自动关闭Dropout、BatchNorm等仅训练阶段生效的算子,和PyTorch切eval模式的行为完全一致。torch.no_grad()等效:默认情况下,不在tf.GradientTape()梯度记录上下文中执行的前向计算不会构建梯度计算图,也不会存储反向传播所需的中间值,运行开销和torch.no_grad()完全一致;如果前向逻辑必须写在梯度上下文内,给不需要回传梯度的输出包裹tf.stop_gradient()即可阻断梯度传播。
注意:不要把
model.trainable = False当成model.eval()的等效操作,前者的作用是冻结模型参数使其不参与训练更新,不会关闭Dropout等训练阶段算子,无法匹配推理模式的行为。
复现代码
如果是纯提取嵌入的推理场景,不需要嵌套额外上下文,直接按如下写法即可完全匹配参考PyTorch代码的行为:
# 加载TF版BERT时保持和PyTorch一致的配置即可,示例: # from transformers import TFBertModel # model = TFBertModel.from_pretrained("bert-base-uncased", output_hidden_states=True) # 前向传入training=False等效model.eval(),默认非梯度上下文自动实现no_grad效果 outputs = model( input_ids=tokens_tensor, token_type_ids=segments_tensors, training=False ) # 和PyTorch逻辑一致,第三个返回值为所有12层的隐藏状态 hidden_states = outputs[2]
如果你的前向逻辑需要写在tf.GradientTape()上下文内(比如嵌入提取是训练流程的一部分,不需要对这部分计算回传梯度),配合tf.stop_gradient()使用即可:
with tf.GradientTape() as tape: outputs = model( input_ids=tokens_tensor, token_type_ids=segments_tensors, training=False ) # 阻断隐藏状态的梯度回传,等效torch.no_grad()的效果 hidden_states = tf.stop_gradient(outputs[2])
内容的提问来源于stack exchange,提问作者user19391408
相关产品推荐
相关产品推荐

