Skip to content

Keras - Real Time Prediction using ResNet Model

本章演示如何使用一个预训练的卷积神经网络(CNN)模型——具体是 ResNet50——进行图像分类。ResNet(残差网络)是一种强大的架构,以其有效训练非常深的网络而闻名。我们将使用在大型 ImageNet 数据集上预训练的 ResNet50 模型,该模型可通过 tf.keras.applications 获取。

工作流程/步骤:

  1. 加载预训练的 ResNet50 模型。
  2. 加载输入图像。
  3. 预处理图像以匹配 ResNet50 期望的格式。
  4. 使用模型预测图像的类别概率。
  5. 将预测结果解码为人类可读的标签。
import tensorflow as tf
from tensorflow.keras.applications.resnet50 import ResNet50, preprocess_input, decode_predictions
from tensorflow.keras.utils import load_img, img_to_array # 更新后的图像工具函数
import numpy as np
# import matplotlib.pyplot as plt # (可选)用于显示图像

步骤 2:加载预训练的 ResNet50 模型

Section titled “步骤 2:加载预训练的 ResNet50 模型”

实例化 ResNet50 模型,指定 weights='imagenet' 以下载并加载预训练的权重。include_top=True 表示我们想要包含最终的分类层。

# 加载在 ImageNet 上预训练的 ResNet50 模型
# 首次使用时将自动下载权重
model = ResNet50(weights='imagenet', include_top=True)

你可以检查模型结构:

# model.summary() # (可选)查看各层

预训练模型期望特定格式的输入图像(尺寸、像素值范围、通道顺序)。ResNet50 通常期望 224x224 像素的图像。

# 指定图像文件路径
# 确保工作目录中有图像文件(例如 'banana.jpg'),
# 或者提供完整路径。
img_path = 'banana.jpg' # 将此修改为你的图像文件
# 加载图像,并将其调整到 ResNet50 期望的目标尺寸
target_size = (224, 224)
img = load_img(img_path, target_size=target_size)
# --- (可选)显示加载的图像 ---
# plt.imshow(img)
# plt.axis('off')
# plt.show()
# -----------------------------------------
# 将 PIL 图像对象转换为 NumPy 数组
img_array = img_to_array(img)
# 扩展维度以创建一个包含 1 张图像的批次
# 形状从 (height, width, channels) 变为 (1, height, width, channels)
img_batch = np.expand_dims(img_array, axis=0)
# 为 ResNet50 预处理图像数据
# 这通常包括特定于模型的均值减去和缩放
img_preprocessed = preprocess_input(img_batch)
print(f"图像已加载并预处理。形状:{img_preprocessed.shape}")

preprocess_input 函数处理 ResNet50 模型所需的特定归一化(例如,将 RGB 转换为 BGR,基于 ImageNet 数据集均值进行零中心化)。

将预处理后的图像批次输入模型的 predict 方法。

predictions = model.predict(img_preprocessed)

predictions 变量现在是一个 NumPy 数组,包含了 1000 个 ImageNet 类别中每个类别的概率得分。对于单张图像输入,形状将是 (1, 1000)。

原始概率不太容易解释。使用 decode_predictions 工具函数可以将这些概率转换为 (class_id, class_name, probability) 元组的列表,通常显示前 N 个预测结果。

# 将预测结果解码为可读的类别名称
# 'top=5' 显示前 5 个最可能的类别
decoded = decode_predictions(predictions, top=5)[0] # 索引 [0] 是因为我们处理的是一个批次包含 1 张图像
print("\n前 5 个预测结果:")
for i, (imagenet_id, label, score) in enumerate(decoded):
print(f"{i + 1}: {label} ({score:.2f})")

输出将取决于你的输入图像。对于典型的香蕉图像,你可能会看到:

图像已加载并预处理。形状:(1, 224, 224, 3)
前 5 个预测结果:
1: banana (0.99)
2: acorn_squash (0.00)
3: cucumber (0.00)
4: pineapple (0.00)
5: butternut_squash (0.00)

这表明模型高度自信图像中包含香蕉。这个过程说明了预训练模型在无需从头开始训练模型的情况下实现快速有效的图像分类的强大之处。