如何提取TF2已训练模型所使用的损失函数本身?
TF2 从keras.Model实例获取损失函数的方法
通用获取方式
情况1:手动compile过的模型
你可以直接访问model.loss属性拿到编译时指定的损失函数:
- 单输出模型该属性直接返回损失函数/损失类实例
- 多输出模型该属性会返回对应每个输出的损失列表或字典,和你compile时传入的结构一致
示例代码:
import tensorflow as tf # 构建示例模型并编译 model = tf.keras.Sequential([ tf.keras.layers.Dense(128, activation='relu'), tf.keras.layers.Dense(10) ]) model.compile( optimizer='adam', loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True) ) # 直接获取损失函数 loss_fn = model.loss # 直接传入标签和logits计算损失 y_true = tf.random.uniform((32,), minval=0, maxval=10, dtype=tf.int32) y_logits = model(tf.random.normal((32, 64))) print(loss_fn(y_true, y_logits)) # 求梯度示例 with tf.GradientTape() as tape: logits = model(tf.random.normal((32, 64)), training=True) loss = loss_fn(y_true, logits) # 对输入求梯度 input_grad = tape.gradient(loss, model.input) # 对模型参数求梯度 param_grads = tape.gradient(loss, model.trainable_variables)
情况2:TF Model Zoo未手动compile的预训练模型
TF Model Zoo的官方模型通常会把损失函数实例绑定在model.loss属性上,直接读取即可。如果部分模型的损失逻辑封装在训练流程中没有直接暴露,你可以去对应模型的构建源码中把损失函数的定义复制出来使用,参数配置和模型训练时保持一致即可。
注意事项
- 不要把
model.loss和model.losses混淆:后者是模型所有正则化损失项的列表,不是你需要的标签和logits对应的预测损失函数 - 使用前确认损失函数的参数要求:比如是否要求输入为logits、是否需要对标签做one-hot编码,避免计算结果不符合预期
- 所有TF原生损失函数和官方实现的自定义损失都默认支持自动微分,直接放在
tf.GradientTape上下文中调用即可正常求梯度
内容的提问来源于stack exchange,提问作者Meysam Sadeghi
相关产品推荐
相关产品推荐

