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

如何使用接收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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 21:36:03