如何使用接收tf.Tensor的preprocess_input预处理tf.data.Dataset
解决方案
方案1:直接用map方法处理tf.data.Dataset
tf.keras.applications.resnet50.preprocess_input原生支持tf.Tensor格式输入,你只需要通过tf.data.Dataset的map接口把预处理逻辑嵌入数据流水线即可:
import tensorflow as tf from tensorflow.keras.applications.resnet50 import preprocess_input # 加载数据集,这里保持你原来的参数即可 train_ds = tf.keras.utils.image_dataset_from_directory( "./your_image_dir", image_size=(224, 224), batch_size=32, label_mode="categorical" ) # 定义预处理函数,同时处理图像和标签 def preprocess_fn(image, label): return preprocess_input(image), label # 映射到数据集,开并行预处理加速 train_ds = train_ds.map(preprocess_fn, num_parallel_calls=tf.data.AUTOTUNE) # 后续可以直接接shuffle、prefetch等常规流水线操作 train_ds = train_ds.shuffle(1000).prefetch(tf.data.AUTOTUNE)
这个方案的优势是预处理逻辑在数据流水线阶段跑在CPU上,不会占用GPU训练资源,适合训练阶段使用。
方案2:把预处理逻辑嵌入模型内部
如果希望预处理和模型绑定、减少部署阶段的额外逻辑,也可以直接把preprocess_input加到模型的输入层之后,不需要改动数据流水线:
from tensorflow.keras import layers # 模型输入层 inputs = layers.Input(shape=(224, 224, 3)) # 直接加入ResNet50预处理逻辑 x = preprocess_input(inputs) # 加载预训练ResNet50 base_model = tf.keras.applications.ResNet50( weights="imagenet", include_top=False, input_tensor=x ) # 后续接你自定义的分类头 x = layers.GlobalAveragePooling2D()(base_model.output) outputs = layers.Dense(你的类别数, activation="softmax")(x) model = tf.keras.Model(inputs=inputs, outputs=outputs)
你也可以手动实现等价的预处理逻辑,避免依赖preprocess_input函数,逻辑完全匹配官方说明:
def custom_resnet_preprocess(inputs): # RGB转BGR:逆序最后一个通道 x = tf.reverse(inputs, axis=[-1]) # 减去ImageNet数据集BGR通道对应的均值,无缩放 mean = tf.constant([103.939, 116.779, 123.68], dtype=tf.float32) return x - mean
两种方案都完全符合ResNet50预训练权重的预处理要求,无需额外调整其他代码。
内容的提问来源于stack exchange,提问作者kymkcay
相关产品推荐
相关产品推荐

