训练非Keras的TensorFlow模型时如何获取中间层(op)输入输出
问题答案
完全可以在训练过程中获取中间层、中间op的输入输出,针对你给出的TF2.8.0+子类Model的写法,有两种最常用的实现方式,都不会破坏模型原有训练逻辑。
方法1:轻改模型call方法,直接挂载中间值(最灵活,适配所有场景)
这个方法适合需要捕获任意中间张量的场景,不管是层输出还是手写op的结果都能拿,调试时用得最多。
你只需要在模型的__init__里加个存中间值的字典,在call方法前向传播的过程中,把需要的输入、中间张量存到这个字典里就行:
class MyModel(Model): def __init__(self): super(MyModel, self).__init__() self.conv1 = Conv2D(32, 3, activation='relu') self.flatten = Flatten() self.d1 = Dense(128, activation='relu') self.d2 = Dense(10) # 初始化中间值存储字典 self.intermediates = {} def call(self, x): x1 = self.conv1(x) # 存conv1的输入、输出 self.intermediates['conv1_input'] = x self.intermediates['conv1_output'] = x1 x2 = self.flatten(x1) self.intermediates['flatten_output'] = x2 x3 = self.d1(x2) self.intermediates['d1_output'] = x3 return self.d2(x3)
训练过程中,每执行完一步前向传播,直接访问model.intermediates就能拿到所有存的中间张量:
- Eager模式下可以直接调用
.numpy()拿到对应的numpy数组值 - 用
tf.function包装训练步的静态图模式下,拿到的是张量引用,在梯度带外调用.numpy()即可取到实际数值 - 注意用完及时清理不需要的中间值引用,避免长期占用显存
方法2:无侵入构造中间值捕获模型(不用改原有模型代码)
如果不想修改已经写好的模型代码,可以利用TF模型权重共享的特性,构造一个专门输出中间张量的副本模型,和原模型权重完全同步,不用重复训练。
注意子类模型是延迟构建的,需要先传一次和输入尺寸一致的dummy张量完成建图,才能拿到各层的输入输出张量:
# 实例化原模型 model = MyModel() # 传入和实际输入形状匹配的哑张量,触发模型结构构建 dummy_x = tf.random.normal((32, 28, 28, 1)) # 以MNIST的batch输入形状为例 _ = model(dummy_x) # 配置需要捕获的张量 capture_outputs = { "conv1_input": model.conv1.input, "conv1_output": model.conv1.output, "x2_flatten": model.flatten.output, "x3_dense1": model.d1.output, "final_logits": model.output } # 构造捕获模型,和原模型共享全部权重 capture_model = tf.keras.Model(inputs=model.input, outputs=capture_outputs)
使用时直接把训练/推理的输入喂给capture_model,返回的字典里就包含所有你配置的中间值,原模型的训练、验证流程完全不受影响。
注意:这个方法只能捕获封装成Layer的层的输入输出,如果你的call方法里存在不属于任何层的独立op(比如手写的
tf.nn.relu、tf.reshape等操作生成的张量),是没法通过层属性拿到的,这种场景用方法1更合适。
可选:结合回调自动记录
如果需要在训练过程中批量记录中间值做分析、可视化,可以自定义Keras回调,在每个batch/epoch结束时自动读取中间值存到日志里,不用手动在训练循环里写读取逻辑。
内容的提问来源于stack exchange,提问作者Ausrada404
相关产品推荐
相关产品推荐

