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

Swin-Transformer-TF使用生成器训练时崩溃,输入尺寸被识别为(None,None,None,3)

Swin-Transformer-TF使用生成器训练时崩溃,输入尺寸被识别为(None,None,None,3)

这个问题我之前踩过坑,核心原因是Swin-Transformer-TF的PatchEmbed层依赖静态形状断言,而普通Python生成器没办法给TensorFlow提供明确的输入静态形状,导致模型误以为输入的高宽都是None,直接触发了断言失败。

问题根源

看报错里的关键代码片段:

# 来自Swin-Transformer-TF的PatchEmbed层实现
(B, H, W, C) = x.get_shape().as_list()
assert H == self.img_size[0] and W == self.img_size[1], \
    f"Input image size ({H}*{W}) doesn't match model ({self.img_size[0]}*{self.img_size[1]})."

这里用get_shape().as_list()获取的是静态形状:当你直接传入numpy数组时,TensorFlow能明确识别出H=224、W=224,断言顺利通过;但用普通生成器时,TensorFlow无法提前推断输入的具体尺寸,静态形状就变成了None,直接触发断言错误。

三种可行解决方案

1. 用tf.data.Dataset包装生成器,明确指定输出形状

这是最省心的方案,不用修改源码,直接给TensorFlow提供明确的形状信息:

def data_generation():
    for i in range(3000):
        yield np.zeros((20,224,224,3)), np.zeros((20,2))

# 用tf.data包装并声明输出形状
dataset = tf.data.Dataset.from_generator(
    data_generation,
    output_signature=(
        tf.TensorSpec(shape=(20, 224, 224, 3), dtype=tf.float32),
        tf.TensorSpec(shape=(20, 2), dtype=tf.float32)
    )
)

# 现在训练就不会报错了
model.fit(dataset, epochs=1, steps_per_epoch=4)

2. 修改PatchEmbed层的断言逻辑,支持动态形状

如果你愿意修改源码,可以把静态断言改成动态检查,让模型支持动态输入尺寸:
找到Swin-Transformer-TF/swintransformer/model.py中的PatchEmbed类,修改它的call方法:

def call(self, x):
    # 改用tf.shape获取动态形状,替代静态形状获取
    B, H, W, C = tf.shape(x)[0], tf.shape(x)[1], tf.shape(x)[2], tf.shape(x)[3]
    # 用TensorFlow的动态断言替代Python原生assert
    tf.debugging.assert_equal(
        H, self.img_size[0],
        f"Input image height {H} doesn't match model {self.img_size[0]}"
    )
    tf.debugging.assert_equal(
        W, self.img_size[1],
        f"Input image width {W} doesn't match model {self.img_size[1]}"
    )
    x = self.proj(x)
    # 用动态计算的尺寸重塑张量
    x = tf.reshape(x, shape=[-1, H//self.patch_size[0] * (W//self.patch_size[0]), self.embed_dim])
    if self.norm is not None:
        x = self.norm(x)
    return x

修改后,模型会在运行时检查输入的实际尺寸,不再依赖静态形状推断。

3. 使用tf.keras.utils.Sequence作为数据生成器

Sequence类是Keras官方推荐的生成器格式,它会自动向模型提供明确的输入形状,兼容性更好:

from tensorflow.keras.utils import Sequence

class DataGenerator(Sequence):
    def __init__(self, batch_size=20, num_batches=3000):
        self.batch_size = batch_size
        self.num_batches = num_batches

    def __len__(self):
        # 返回训练总步数
        return self.num_batches

    def __getitem__(self, idx):
        # 生成单个批次的数据
        x = np.zeros((self.batch_size, 224, 224, 3))
        y = np.zeros((self.batch_size, 2))
        return x, y

# 创建生成器实例
gen = DataGenerator()
model.fit(gen, epochs=1, steps_per_epoch=4)

验证说明

这三种方案都能解决你的问题,其中方案1最简洁,不需要修改源码;方案2适合需要自定义动态尺寸检查的场景;方案3则是Keras生态下的标准做法,扩展性更强。

备注:内容来源于stack exchange,提问作者Naomi Fridman

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.21 14:17:57