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
相关产品推荐
相关产品推荐

