You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.07 17:05:18