1. 项目概述与需求分析
1.1 为什么做运动鞋识别
先说个场景。我平时喜欢逛街的时候看鞋,尤其是运动鞋。但问题在于,很多潮流鞋款、联名款、复刻款,光靠肉眼分辨极其困难。有一次我朋友拿着一双鞋问我"这是不是AJ1芝加哥",我看了半天也没敢确定。后来我想到,这不是一个典型的图像分类问题吗?干脆自己训练一个运动鞋识别模型,用TensorFlow来做。这个项目的核心目标很直接:输入一张运动鞋照片,模型输出这双鞋的型号或类别。考虑到实际使用场景,可能需要识别常见的耐克、阿迪、彪马、New Balance等品牌,甚至具体到AJ1、Air Max、Yeezy这样的型号。
这个项目适合谁?如果你是刚接触计算机视觉的开发者,想用TensorFlow练手,这个项目非常合适。因为运动鞋图像的类别相对固定,背景可以控制,数据采集相对容易(网上素材极多),而且模型的可解释性好——识别对了就是对了,错了也容易找到原因。相比于做那种"猫狗分类"的经典demo,运动鞋识别更有趣味性,也更贴近真实的产品需求。另外,如果你的工作涉及电商、二手交易平台,这类识别模型其实有实实在在的落地价值,比如自动识别商品类目、辅助鉴定真伪等场景。
整个项目我基于TensorFlow 2.x版本实现,结合了迁移学习、数据增强、模型导出等技术点。下面把所有细节拆开讲,从环境搭建到最终推理一条龙,希望能帮你复现出属于自己的运动鞋识别模型。
1.2 技术选型:为什么是TensorFlow
说实话,2024年PyTorch在学术界确实很流行,很多新论文的官方代码都是PyTorch写的。但TensorFlow的优势在于工程化和部署生态成熟。我的项目要求是快速训练、导出模型、然后能在手机或嵌入式设备上跑推理,TensorFlow的TFLite和TF Serving在这一块可以说是无缝衔接。TensorFlow 2.x之后的Eager Execution(动态图模式)让模型调试也变得直观,不再有早期那种"先建图再会话"的割裂感。
另外一个现实原因是,我的数据管道是用tf.data搭建的,训练和部署可以完全在一个框架内闭环。如果你在做一个正经小项目,而不是单纯跑跑别人的demo,TensorFlow的Keras API加上TensorBoard可视化,配合TFRecord数据格式,整个流程非常顺手。PyTorch当然也好,但我不想为了赶学术时髦而增加额外的工作量——工程落地才是这个项目的重点。
这里补充一句,TensorFlow 2.18是目前比较新的稳定版本,安装体验比早期版本好太多,对Python 3.11/3.12的支持也完善了。各位在装环境的时候千万别装1.x的老古董,直接用最新稳定版即可。
需要模型API调用? 免费领10W Token,多模型网关一键接入 Claude、DeepSeek 等主流模型。
2. 环境准备与依赖安装
2.1 TensorFlow 2.18安装细节
安装TensorFlow看起来简单,就是pip install tensorflow一条命令。但如果你对GPU版本有要求,这里有几个坑值得说一说。我项目的训练环境是Windows 11 + NVIDIA RTX 3060,CUDA 12.3,cuDNN 8.9。TensorFlow 2.18版本对应的CUDA和cuDNN版本终于统一到了CUDA 12系列,这比之前2.10那会儿要求CUDA 11.2舒服多了。
注意:TensorFlow 2.18默认安装时会一并安装numpy、keras等依赖。如果你机器上已经装了别的版本的numpy,强烈建议在虚拟环境里装,避免依赖冲突。
我的安装命令(在conda虚拟环境里):
bash复制conda create -n sneaker_tf python=3.11
conda activate sneaker_tf
pip install tensorflow==2.18
pip install tensorflow-gpu==2.18 # 如果你之前装了CPU版,这一条会提示冲突,需要先卸载
我这里直接用CPU版其实也能跑,因为图像分类的小数据集训练量不大。但如果你的数据集有数万张图,那就老老实实装GPU版。我实测下来,RTX 3060训练一个EfficientNetB0的模型,只需CPU的十分之一时间。
安装后务必验证一下版本和GPU是否可用:
python复制import tensorflow as tf
print(tf.__version__) # 输出2.18.0
print(tf.config.list_physical_devices('GPU')) # 查看GPU是否被识别
如果你看到GPU列表为空,大概率是CUDA、cuDNN的版本不匹配,或者是没安装对应的NVIDIA驱动。这一块我在第6章专门讲排查方法。
2.2 目录结构与项目初始化
一个好的项目结构能让你少走弯路。我的目录如下:
code复制sneaker-recognition/
├── data/
│ ├── raw/ # 原始图片,按类别分文件夹
│ ├── processed/ # 预处理后图片,统一尺寸
│ └── tfrecords/ # TFRecord文件
├── models/ # 保存训练好的模型
├── notebooks/ # 实验用的Jupyter笔记本
├── scripts/
│ ├── prepare_data.py # 数据预处理脚本
│ ├── train.py # 训练脚本
│ ├── evaluate.py # 评估脚本
│ └── inference.py # 推理脚本
├── requirements.txt
└── README.md
好代码的关键在于可复现、可维护。我把数据准备、训练、评估分成独立模块,这样改数据增强策略就不用动训练代码,换模型也容易。
3. 数据准备与预处理
3.1 数据来源与标注策略
运动鞋图片怎么来?我用了两个渠道:
- 从电商平台公开图片爬取(注意版权和合规问题,这里只是个人学习研究用途)
- 从已有的开源数据集,比如斯坦福的鞋类数据集、Fashion MNIST这类数据里做扩充
我最终选了8个类别的运动鞋:Nike Air Max 270、Adidas Ultraboost、New Balance 574、Puma Suede、Reebok Classic、Converse Chuck Taylor(严格说不是运动鞋但属于球鞋类)、Vans Old Skool、Yeezy Boost 350V2。每个类别各收集400张图,总共3200张。这个数据量对于迁移学习来说足够入门了。
标注方面,就是按文件夹分好类,文件夹名用英文标识符,比如nike_air_max_270、adidas_ultraboost。用tf.keras.preprocessing.image_dataset_from_directory读取时,它会自动根据子文件夹名生成标签。标签就是0到7的整数索引。
3.2 数据清洗与增强
原始图片千奇百怪:有的带水印,有的背景复杂,有的分辨率极低。我的处理原则:
- 剔除模糊、严重遮挡、多只鞋同框的图片
- 统一裁剪为224x224分辨率,这对ImageNet预训练模型的输入尺寸是标准的
- 归一化到[0,1]区间,即像素值除以255
数据增强我用的是tf.keras自带的预处理层:
python复制data_augmentation = tf.keras.Sequential([
tf.keras.layers.RandomFlip("horizontal"),
tf.keras.layers.RandomRotation(0.1),
tf.keras.layers.RandomZoom(0.1),
tf.keras.layers.RandomContrast(0.1),
])
增强的目的是防止过拟合。运动鞋图片有一个特点——鞋的朝向相对固定,但角度、光线的变化很大。水平翻转对运动鞋识别有效,因为左右反转并不改变鞋型。但垂直翻转就不要用,因为倒着的鞋看起来很不自然,反而会误导模型。
实操心得:对于运动鞋这类目标,我更推荐在训练阶段做增强,推理阶段不做增强。原因很简单,推理时用户拍的图片可能千奇百怪,但你不能把一张图片翻转成四张来投票吧?当然,如果你需要更精确的预测,可以在推理时做多尺度预测,即把图片缩小放大分别预测再取平均,但这个对嵌入式部署不友好。
训练集、验证集、测试集按照8:1:1划分。由于数据量不大,我没有用TFRecord,直接用image_dataset_from_directory:
python复制train_ds = tf.keras.preprocessing.image_dataset_from_directory(
'data/processed',
validation_split=0.2,
subset='training',
seed=123,
image_size=(224, 224),
batch_size=32,
label_mode='int'
)
这样索引比较方便,适合快速上手。
4. 模型设计与训练
4.1 从头搭建CNN vs 迁移学习
运动鞋分类是一个相对细粒度的视觉任务,不同鞋款之间的差异可能只在鞋面线条、鞋底纹路、boost颗粒等细节上。用VGG16这种经典结构加上ImageNet预训练权重,效果会比从零训练好得多。我实验过从零开始训练一个简单的LeNet变体,准确率只有78%,跟瞎猜差不多。而用EfficientNetB0做特征提取器,准确率能到93%。
这里我推荐两个迁移学习策略:
- 把预训练模型当成特征提取器:冻结卷积基,只训练顶层的全连接分类器。适合数据量少、训练时间短的情况。
- 微调(Fine-tune):在冻结预训练模型训练几轮后,解冻部分层,用较小的学习率继续训练。适合追求更高准确率的场景。
我的方案是先用策略1跑10个epoch,再解冻最后几层,用策略2微调5个epoch。最终验证准确率能达到96%。这主要是因为ImageNet预训练模型已经掌握了大量底层的纹理、边缘特征,而运动鞋的纹理边缘和自然图像有一定的共性。微调阶段让模型适应鞋子特有的颜色和形状,效果顿时提升。
4.2 搭建模型
我用的是EfficientNetB0,这个模型参数量不大,推理速度快,很适合部署在嵌入式设备上。代码如下:
python复制def build_model(num_classes=8):
base_model = tf.keras.applications.EfficientNetB0(
include_top=False,
weights='imagenet',
input_shape=(224, 224, 3)
)
base_model.trainable = False # 先冻结
inputs = tf.keras.Input(shape=(224, 224, 3))
x = data_augmentation(inputs) # 增强在模型内部做
x = tf.keras.applications.efficientnet.preprocess_input(x)
x = base_model(x, training=False)
x = tf.keras.layers.GlobalAveragePooling2D()(x)
x = tf.keras.layers.Dropout(0.2)(x)
outputs = tf.keras.layers.Dense(num_classes, activation='softmax')(x)
model = tf.keras.Model(inputs, outputs)
return model
注意几个细节:
preprocess_input会把像素值从[0,255]映射到[-1,1],这个和手写归一化等价,但直接调用更标准。GlobalAveragePooling2D替代了Flatten加全连接,能够大幅减少参数量,同时防止过拟合。Dropout(0.2)是防止全连接层过拟合的经典手段。
4.3 训练参数与优化器选择
训练参数如下:
- 优化器:Adam,初始学习率0.001,微调阶段降到0.0001
- 损失函数:SparseCategoricalCrossentropy(因为标签是整数索引)
- 批次大小:32
- Epochs:特征提取阶段10轮,微调阶段5轮
- 学习率调度:ReduceLROnPlateau,监控验证损失,如果3轮不下降则学习率乘以0.5
这里说一下优化器选择的逻辑。Adam在分类任务上是默认选项,它结合了Momentum和RMSProp的优点,收敛速度快,对学习率不那么敏感。SGD+Momentum也很好,但需要手动调整学习率和动量参数,新手容易调崩。所以建议直接用Adam。
另外,我在训练时加了EarlyStopping:
python复制callbacks = [
tf.keras.callbacks.EarlyStopping(monitor='val_loss', patience=5, restore_best_weights=True),
tf.keras.callbacks.ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=3),
tf.keras.callbacks.ModelCheckpoint('models/sneaker_model.h5', save_best_only=True)
]
restore_best_weights非常关键。否则如果第8轮出现了过拟合,模型最后保存的权重可能是第10轮的,而不是验证集上最好的那一轮。踩过坑的人都知道,这个参数能救你的模型。
4.4 训练过程解析
我用TensorBoard记录训练曲线。特征提取阶段前3轮准确率就从50%上升到85%左右,说明预训练模型的特征非常有效。验证准确率在91%左右波动,有些过拟合迹象(训练准确率接近100%),这就是需要增强和Dropout的原因。
微调阶段,我解冻了EfficientNetB0的最后20层,统称base_model.trainable = True,但需要配合更小的学习率,否则容易破坏预训练学到的良好的低级特征。微调后验证准确率稳定在96~97%。对于8类运动鞋来说,这个准确率已经够用了。
这里要特别说明,微调阶段不是解冻所有层。设计者的经验是,越靠近输入的层学习到的是边缘、颜色等通用特征,不适合大规模更新。所以我用循环选择解冻范围:
python复制base_model.trainable = True
for layer in base_model.layers[:200]:
layer.trainable = False
EfficientNetB0一共大概200多层,我只让最后几十层可训练,这样既不会丢失通用特征,又能适应运动鞋的特定细节。
5. 评估、推理与导出
5.1 评估指标的选择
准确率不等于一切。对于8类分类,还要看每一类的精确率、召回率和F1分数。我计算了混淆矩阵,发现最容易混淆的是Adidas Ultraboost和Yeezy Boost 350V2,因为它们都是网面鞋型,鞋底偏厚,从侧上方角度拍确实很像。
用sklearn生成混淆矩阵和分类报告:
python复制from sklearn.metrics import classification_report, confusion_matrix
y_true = ...
y_pred = model.predict(test_ds).argmax(axis=-1)
print(classification_report(y_true, y_pred))
好的消息是,每类平均召回率都在94%以上,这说明模型没有明显的短板。如果你想进一步提高,可以考虑增加这两类容易混淆的图片数量,或者在数据增强中加入色彩抖动,因为实际场景中这两类鞋最明显的差异在颜色和鞋面纹路。
5.2 单图推理与模型导出
训练完成后,我写了推理脚本:
python复制def predict_sneaker(image_path, model_path='models/sneaker_model.h5'):
model = tf.keras.models.load_model(model_path)
img = tf.keras.preprocessing.image.load_img(image_path, target_size=(224, 224))
img_array = tf.keras.preprocessing.image.img_to_array(img)
img_array = tf.expand_dims(img_array, 0)
img_array = tf.keras.applications.efficientnet.preprocess_input(img_array)
preds = model.predict(img_array)
idx = tf.argmax(preds[0]).numpy()
class_names = list(train_ds.class_names)
confidence = tf.nn.softmax(preds[0])[idx].numpy()
return class_names[idx], float(confidence)
这个脚本可以直接在终端运行,也可以封装成Flask的API接口,用户上传图片即可返回预测结果。
导出TFLite模型以便在手机端使用:
python复制converter = tf.lite.TFLiteConverter.from_keras_model(model)
converter.optimizations = [tf.lite.Optimize.DEFAULT]
tflite_model = converter.convert()
with open('models/sneaker_model.tflite', 'wb') as f:
f.write(tflite_model)
这里使用了默认优化,即量化权重到float16或int8,模型体积会缩小到原来的四分之一,推理速度加快。对于运动鞋识别这种对实时性要求不高的场景,量化的影响可以忽略。
6. 常见问题与排查技巧实录
6.1 TensorFlow安装与GPU问题
问题1:pip install tensorflow 安装后,import时报错找不到某个DLL(Windows)。
原因大多是缺少Visual C++ Redistributable。装了最新的VC++运行库后重启就好了。
问题2:GPU无法使用
检查NVIDIA驱动版本,确认CUDA 12.3与cuDNN 8.9是否已安装,且环境变量里包含路径。Windows下还要把bin目录加入PATH,可能还需要LD_LIBRARY_PATH(Linux)或PATH(Windows)。
问题3:某天突然报Could not create cudnn handle: CUDNN_STATUS_INTERNAL_ERROR
这种通常是显存不足或cuDNN缓存问题。试一下在代码开头加:
python复制gpus = tf.config.list_physical_devices('GPU')
if gpus:
try:
tf.config.experimental.set_memory_growth(gpus[0], True)
except RuntimeError as e:
print(e)
这样TensorFlow不会一次性占满显存,而是按需分配。
6.2 数据处理相关的问题
图片读取时遇到不可解码的文件,通常是因为某些图片损坏或格式不是标准的JPEG/PNG。我写了一个简单的清理脚本,遍历所有图片,用PIL打开,打不开就删除或记录路径。
python复制from PIL import Image
import os
root = 'data/raw'
invalid = []
for label in os.listdir(root):
label_path = os.path.join(root, label)
if not os.path.isdir(label_path):
continue
for img_name in os.listdir(label_path):
img_path = os.path.join(label_path, img_name)
try:
with Image.open(img_path) as im:
im.load()
except Exception:
invalid.append(img_path)
print(f'发现 {len(invalid)} 个损坏文件')
数据增强时注意不要用太强的旋转,因为运动鞋的朝向通常固定,旋转角度超过30度会让鞋变得面目全非。建议rotation的范围控制在±15度以内。
6.3 与PyTorch对比的几个坑
我在另一个项目里用PyTorch做类似任务,对比下来,TensorFlow这边有几个独有的坑:
-
数据管道。TensorFlow的
tf.data性能好,但调试起来没有PyTorch的DataLoader直观。打印一个batch的数据,PyTorch可以直接看到张量值,tf.data却要在迭代器里next(iter(ds)),类型转换还要带.numpy()。 -
模型加载。TensorFlow的
.h5文件需要完全一致的模型定义才能加载,所以最好把模型定义和训练逻辑分开。PyTorch的.pth则保存的是OrderedDict,也是依赖模型结构。这一点两者差不多,但TensorFlow的SavedModel格式携带了完整的图结构,部署时更友好。 -
Keras的Functional API非常直观。在搭建复杂模型时比PyTorch的
nn.Module少很多模板代码。但PyTorch的灵活度更高,比如自定义循环、中间层输出等。如果你要频繁做实验探索,PyTorch更顺;如果快速落地,TensorFlow更稳。
6.4 训练过拟合的排查
如果训练准确率持续上升,验证准确率上不去,说明过拟合了。我遇到过几次,典型的处理方案:
- 增加数据增强的强度,比如加入随机颜色抖动、亮度变化
- 降低模型复杂度,比如减少Dropout、减少全连接层
- 增加更多数据,真实数据永远是王道
- 提前停止,用EarlyStopping在验证loss上升时及时终止
运动鞋图片有一个特点,很多图片背景是纯色(白墙、地板),模型容易学习"背景特征"而不是"鞋的特征"。这种时候建议在预处理时把图片进行裁剪,使鞋占据画面主要部分,或者使用背景替换技术。不过背景替换比较复杂,我暂时用随机裁剪来模拟不同的构图。
7. 项目扩展思路与个人体会
这个运动鞋识别项目做到这里,核心流程已经跑通了。如果你想让项目更有实战价值,可以考虑这几个方向:
第一,把分类任务改成目标检测。比如用TensorFlow Object Detection API或者EfficientDet,检测图片中的鞋并定位其位置。在鞋类电商场景中,用户上传的图片往往包含多个物品,甚至鞋子和人的脚一起出现在画面中。训练一个检测器,先找到鞋的边界框,再做分类,整个系统会更鲁棒。
第二,加入细粒度识别。运动鞋领域,很多人关心的是同一品牌下的不同款式。比如耐克的Air Force 1和Air Jordan 1,外形长得非常像,需要模型关注鞋面、鞋帮、鞋底纹路等细节。这种情况下,可以考虑注意力机制或者对比学习,进一步提高识别精度。
第三,部署到移动端或嵌入式设备。把TFLite模型集成到Android或iOS应用中,用户拍一张照片,立即返回识别结果。这是目前我认为最有趣的扩展方向。
最后分享一下我踩了几次坑之后的感受。做图像分类项目,真正难的地方往往不是模型结构,而是数据。第一次训练时,我为了图省事直接从网上抓了一堆图片,没有清洗,导致模型准确率卡在80%上不去。后来花了一整天清理数据,把模糊的、带水印的、背景杂乱的图片全部过滤掉,准确率立刻提升到95%以上。数据质量决定了模型上限,这句话在这个项目里体现得淋漓尽致。
另外,关于TensorFlow和PyTorch的选型,我的态度很直接:如果你是做学术研究或者经常要复现论文,PyTorch更迎合潮流;但如果你要做产品落地,想把模型快速部署到各类平台,TensorFlow的生态还是有一定优势的。反正技术栈不是终点,能解决问题、能维护、能上线才是重点。运动鞋识别这个项目用TensorFlow走通之后,我对整个图像分类的工程流程已经烂熟于胸,下次再遇到类似的分类任务,基本可以一天之内复现出结果。
