如何在Keras子类化模型中固定Batch Size?
在Keras子类化模型中固定Batch Size的方法
在Keras子类化模型中,你可以通过以下两种常见方式固定batch size:
方法一:在__init__中定义带batch_size的输入层
直接在模型初始化方法里创建指定batch_size的Input层,后续层基于该输入层构建模型结构。
示例代码:
import tensorflow as tf from tensorflow import keras class CustomModel(keras.Model): def __init__(self): super().__init__() # 定义固定batch_size为1的输入层 self.input_layer = keras.Input(shape=(64, 64, 3), batch_size=1) self.conv_transpose = keras.layers.Conv2DTranspose(3, 3, strides=2, padding="same", activation="relu") def call(self, inputs): x = self.conv_transpose(inputs) return x # 实例化模型并基于预定义输入层确定结构 model = CustomModel() _ = model(model.input_layer)
方法二:在build方法中指定包含batch size的输入shape
在模型的build方法里,强制定义包含固定batch size的输入形状,让模型构建时锁定batch维度。
示例代码:
import tensorflow as tf from tensorflow import keras class CustomModel(keras.Model): def __init__(self): super().__init__() self.conv_transpose = keras.layers.Conv2DTranspose(3, 3, strides=2, padding="same", activation="relu") def build(self, input_shape): # 覆盖输入shape,固定batch size为1 fixed_input_shape = (1,) + input_shape[1:] super().build(fixed_input_shape) def call(self, inputs): # 可选:添加断言校验输入batch size是否符合要求 assert inputs.shape[0] == 1, f"输入batch size必须为1,当前为{inputs.shape[0]}" x = self.conv_transpose(inputs) return x # 实例化模型并触发构建 model = CustomModel() model.build((None, 64, 64, 3))
你也可以仅在call方法中添加batch size的校验逻辑,作为对输入的约束补充。
内容的提问来源于stack exchange,提问作者YeongHwa Jin
相关产品推荐
相关产品推荐

