YOLOX实战:从零构建高精度目标检测模型的避坑全流程
如果你正在Ubuntu系统上尝试用PyTorch训练自己的YOLOX模型,大概率已经体会过那种“明明照着教程走,却总在奇怪的地方卡住”的挫败感。环境配置报错、数据集格式混乱、训练中途崩溃、模型加载失败……这些坑我一个都没落下,全都踩了一遍。
今天这份指南,就是我在实际项目中用YOLOX完成多个工业检测任务后,整理出的完整避坑手册。我不会简单重复官方文档的内容,而是聚焦于那些官方没细说、但实践中一定会遇到的关键问题。无论你是刚接触目标检测的新手,还是想从其他框架迁移到YOLOX的开发者,这篇文章都能帮你节省大量调试时间。
1. 环境配置:不只是安装依赖那么简单
很多人以为环境配置就是pip install几行命令的事,结果往往在第一个demo就跑不起来。YOLOX的环境配置有几个特别容易出错的点,尤其是PyTorch、CUDA和apex的版本兼容问题。
1.1 版本匹配:PyTorch与CUDA的“婚姻关系”
YOLOX官方推荐PyTorch 1.7+,但并不意味着版本越高越好。我最初用PyTorch 1.9就遇到了CUDA扩展编译失败的问题。关键是要确保PyTorch、CUDA、cuDNN三者版本匹配。
推荐组合方案:
- Ubuntu 18.04 + CUDA 11.0 + cuDNN 8.0.5 + PyTorch 1.7.1
- Ubuntu 20.04 + CUDA 11.1 + cuDNN 8.0.5 + PyTorch 1.8.0
如果你已经安装了其他版本,可以通过conda快速切换:
# 创建新环境
conda create -n yolox python=3.8
conda activate yolox
# 安装指定版本的PyTorch
conda install pytorch==1.7.1 torchvision==0.8.2 torchaudio==0.7.2 cudatoolkit=11.0 -c pytorch
注意:不要盲目使用
pip install torch,这通常会安装最新版,可能与你的CUDA版本不兼容。先通过nvcc --version和nvidia-smi确认CUDA版本,再选择对应的PyTorch版本。
1.2 Apex安装:混合精度训练的关键
YOLOX默认开启混合精度训练(FP16),这需要NVIDIA的apex库。但apex的安装是最大的坑之一,特别是CUDA版本不匹配时的编译错误。
避坑安装法:
# 1. 先卸载可能存在的旧版本
pip uninstall apex -y
# 2. 克隆源码(推荐指定分支)
git clone https://github.com/NVIDIA/apex
cd apex
# 3. 如果遇到CUDA版本不匹配的报错,修改setup.py
# 找到check_cuda_torch_binary_vs_bare_metal函数
# 在函数开头直接添加 return 语句,跳过版本检查
# 修改后保存
# 4. 安装
pip install -v --no-cache-dir --global-option="--cpp_ext" --global-option="--cuda_ext" ./
如果还是失败,可以尝试不编译CUDA扩展的简化安装:
pip install -v --no-cache-dir ./
不过这样会失去一些优化,训练速度可能受影响。
1.3 验证环境:用demo测试的正确姿势
环境装好后,很多人直接跑官方demo,没报错就以为成功了。其实demo可能因为各种原因静默失败。正确的验证流程应该是:
# test_environment.py
import torch
import torchvision
import sys
print(f"PyTorch版本: {torch.__version__}")
print(f"CUDA可用: {torch.cuda.is_available()}")
print(f"CUDA版本: {torch.version.cuda}")
print(f"cuDNN版本: {torch.backends.cudnn.version()}")
# 测试张量计算
if torch.cuda.is_available():
x = torch.randn(3, 640, 640).cuda()
y = torch.randn(3, 640, 640).cuda()
z = torch.matmul(x, y)
print("GPU计算测试通过")
else:
print("警告:CUDA不可用,训练将非常缓慢")
然后再下载预训练权重运行demo:
# 下载YOLOX-S权重
wget https://github.com/Megvii-BaseDetection/YOLOX/releases/download/0.1.1rc0/yolox_s.pth
# 运行demo(添加--debug参数查看详细输出)
python tools/demo.py image \
-f exps/default/yolox_s.py \
-c yolox_s.pth \
--path assets/dog.jpg \
--conf 0.3 \
--nms 0.65 \
--tsize 640 \
--save_result \
--device gpu \
--debug
如果输出图片保存在YOLOX_outputs/demo/vis_res/目录下,并且有检测框,说明环境基本正常。
2. 数据集制作:VOC格式的现代实践
YOLOX支持COCO和VOC格式,对于自定义数据集,VOC格式更简单。但官方文档对数据集准备的描述过于简略,这里给出一个工业级的制作流程。
2.1 标注工具选择与格式转换
现在很少有人直接写XML了,通常用标注工具生成JSON,再转换为VOC格式。我推荐两种工作流:
方案A:LabelImg + 自动转换
# 安装LabelImg
pip install labelImg
labelImg # 启动图形界面
# 标注后得到XML文件,使用官方转换脚本
python tools/datasets/voc2coco.py \
--voc_dir ./data/VOCdevkit/VOC2007 \
--output_dir ./data/COCO \
--split trainval
方案B:Roboflow在线标注 + 下载 对于团队协作,Roboflow更高效。标注后可以直接导出为YOLOX所需的格式:
- 在Roboflow上创建项目并上传图片
- 使用Web工具标注
- 导出时选择"VOC XML"格式
- 下载后的结构直接可用
2.2 自动化数据集处理脚本
手动处理数据集容易出错,我写了一个健壮的Python脚本,处理各种边缘情况:
# prepare_dataset.py
import os
import xml.etree.ElementTree as ET
import json
from pathlib import Path
import shutil
from sklearn.model_selection import train_test_split
import cv2
class VOCDatasetPreparer:
def __init__(self, raw_data_dir, output_dir="./data/VOCdevkit/VOC2007"):
self.raw_dir = Path(raw_data_dir)
self.output_dir = Path(output_dir)
self.classes = [] # 会自动从标注中提取
def discover_classes(self):
"""自动发现所有类别"""
classes_set = set()
for xml_file in self.raw_dir.glob("*.xml"):
tree = ET.parse(xml_file)
root = tree.getroot()
for obj in root.findall('object'):
class_name = obj.find('name').text
classes_set.add(class_name)
self.classes = sorted(list(classes_set))
print(f"发现 {len(self.classes)} 个类别: {self.classes}")
return self.classes
def validate_annotations(self):
"""验证标注的完整性"""
issues = []
for xml_file in self.raw_dir.glob("*.xml"):
# 检查对应的图片是否存在
img_name = xml_file.stem + ".jpg"
img_path = self.raw_dir / img_name
if not img_path.exists():
issues.append(f"图片缺失: {img_name}")
continue
# 检查标注是否为空
tree = ET.parse(xml_file)
root = tree.getroot()
objects = root.findall('object')
if len(objects) == 0:
issues.append(f"空标注: {xml_file.name}")
if issues:
print("发现以下问题:")
for issue in issues:
print(f" - {issue}")
return False
return True
def split_dataset(self, test_size=0.2, val_size=0.1, random_seed=42):
"""智能划分数据集"""
all_files = [f.stem for f in self.raw_dir.glob("*.xml")]
# 先分测试集
trainval_files, test_files = train_test_split(
all_files, test_size=test_size, random_state=random_seed
)
# 再从训练集中分验证集
actual_val_size = val_size / (1 - test_size)
train_files, val_files = train_test_split(
trainval_files, test_size=actual_val_size, random_state=random_seed
)
# 保存划分结果
splits_dir = self.output_dir / "ImageSets" / "Main"
splits_dir.mkdir(parents=True, exist_ok=True)
for split_name, files in [("train", train_files),
("val", val_files),
("test", test_files),
("trainval", trainval_files)]:
with open(splits_dir / f"{split_name}.txt", "w") as f:
f.write("\n".join(files))
print(f"数据集划分完成:")
print(f" 训练集: {len(train_files)} 个样本")
print(f" 验证集: {len(val_files)} 个样本")
print(f" 测试集: {len(test_files)} 个样本")
return train_files, val_files, test_files
def prepare(self):
"""执行完整的准备流程"""
print("开始准备数据集...")
# 1. 创建目录结构
(self.output_dir / "Annotations").mkdir(parents=True, exist_ok=True)
(self.output_dir / "JPEGImages").mkdir(parents=True, exist_ok=True)
# 2. 验证数据
if not self.validate_annotations():
print("数据验证失败,请先修复问题")
return
# 3. 发现类别
self.discover_classes()
# 4. 复制文件
print("复制文件...")
for xml_file in self.raw_dir.glob("*.xml"):
# 复制XML
shutil.copy2(xml_file, self.output_dir / "Annotations" / xml_file.name)
# 复制图片
img_file = xml_file.with_suffix(".jpg")
if img_file.exists():
shutil.copy2(img_file, self.output_dir / "JPEGImages" / img_file.name)
# 5. 划分数据集
self.split_dataset()
# 6. 保存类别文件
with open(self.output_dir.parent / "classes.txt", "w") as f:
f.write("\n".join(self.classes))
print("数据集准备完成!")
print(f"输出目录: {self.output_dir.absolute()}")
# 使用示例
if __name__ == "__main__":
preparer = VOCDatasetPreparer(
raw_data_dir="./raw_data", # 原始标注和图片目录
output_dir="./data/VOCdevkit/VOC2007"
)
preparer.prepare()
这个脚本会自动处理类别发现、数据验证、智能划分,比手动操作可靠得多。
2.3 数据集结构检查清单
完成数据集准备后,用这个检查清单验证:
# 检查目录结构
tree data/VOCdevkit -L 3
# 应该看到类似结构
data/VOCdevkit/
└── VOC2007
├── Annotations
│ ├── 000001.xml
│ └── 000002.xml
├── ImageSets
│ └── Main
│ ├── train.txt
│ ├── val.txt
│ ├── test.txt
│ └── trainval.txt
└── JPEGImages
├── 000001.jpg
└── 000002.jpg
# 检查标注数量匹配
ls data/VOCdevkit/VOC2007/Annotations/*.xml | wc -l
ls data/VOCdevkit/VOC2007/JPEGImages/*.jpg | wc -l
# 检查划分文件
cat data/VOCdevkit/VOC2007/ImageSets/Main/train.txt | head -5
3. 配置文件修改:那些容易忽略的细节
YOLOX的配置文件系统比较灵活,但也容易配置错误。以下是必须修改的几个关键文件。
3.1 类别定义与映射
首先修改类别文件,注意有两个地方需要同步修改:
文件1: yolox/data/datasets/voc_classes.py
# 原内容
VOC_CLASSES = (
"aeroplane", "bicycle", "bird", "boat", "bottle", "bus", "car",
"cat", "chair", "cow", "diningtable", "dog", "horse", "motorbike",
"person", "pottedplant", "sheep", "sofa", "train", "tvmonitor",
)
# 修改为你的类别
MY_CLASSES = (
"defect_a", # 你的第一个类别
"defect_b", # 你的第二个类别
# ... 最多不要超过20个类别
)
# 同时修改这个变量
VOC_CLASSES = MY_CLASSES
文件2: exps/example/yolox_voc/yolox_voc_s.py(或其他你使用的exp文件)
class Exp(MyExp):
def __init__(self):
super(Exp, self).__init__()
self.num_classes = 2 # 修改为你的类别数
self.depth = 0.33
self.width = 0.50
self.exp_name = os.path.split(os.path.realpath(__file__))[1].split(".")[0]
# 数据集路径 - 修改这里!
self.data_dir = "./data/VOCdevkit"
self.train_ann = "voc_2007_trainval"
self.val_ann = "voc_2007_test"
# 输入尺寸
self.input_size = (640, 640)
self.test_size = (640, 640)
# 训练轮数 - 根据数据集大小调整
self.max_epoch = 300
self.no_aug_epochs = 15
# 学习率 - 小数据集可以调小
self.basic_lr_per_img = 0.01 / 64.0
3.2 修复VOC数据加载器的bug
YOLOX的VOC数据加载器有个已知bug,在yolox/data/datasets/voc.py中:
# 找到这行代码(大约在第100行)
image_sets = [('2007', 'trainval')] # 原代码
# 修改为
image_sets = [('2007', 'trainval')] # 如果只有2007数据集
# 或者如果你有2007和2012的数据
# image_sets = [('2007', 'trainval'), ('2012', 'trainval')]
还有一个常见问题是在load_anno方法中,如果遇到FileNotFoundError,需要检查路径拼接逻辑。我建议添加一些调试输出:
def load_anno(self, index):
img_info = self.img_infos[index]
filename = img_info['file_name']
# 添加调试信息
xml_path = os.path.join(self.root, "Annotations", filename + ".xml")
if not os.path.exists(xml_path):
print(f"警告: XML文件不存在: {xml_path}")
print(f"当前目录: {os.getcwd()}")
print(f"root路径: {self.root}")
# ... 其余代码不变
3.3 自定义数据增强策略
YOLOX默认的数据增强很强,但对于小数据集或特殊场景可能需要调整。在exp文件中可以修改:
class Exp(MyExp):
def __init__(self):
super(Exp, self).__init__()
# ... 其他配置
# 数据增强配置
self.degrees = 10.0 # 随机旋转角度
self.translate = 0.1 # 平移
self.scale = (0.5, 1.5) # 缩放范围
self.mosaic_scale = (0.5, 1.5) # Mosaic增强的缩放
self.shear = 2.0 # 剪切变换
self.perspective = 0.0 # 透视变换
self.enable_mixup = True # 是否启用MixUp
def get_data_loader(self, batch_size, is_distributed, no_aug=False):
# 可以在这里覆盖数据加载逻辑
from yolox.data import TrainTransform
dataset = VOCDetection(
data_dir=self.data_dir,
image_sets=[('2007', 'trainval')],
img_size=self.input_size,
preproc=TrainTransform(
rgb_means=(0.485, 0.456, 0.406),
std=(0.229, 0.224, 0.225),
max_labels=50, # 最大标注数,根据你的数据集调整
),
)
# ... 其余代码
4. 训练策略:从预训练到微调的最佳实践
训练阶段是最容易出问题的环节,尤其是内存溢出、梯度爆炸、过拟合等问题。
4.1 预训练权重的正确使用
YOLOX提供了在COCO上预训练的权重,但直接加载可能会遇到维度不匹配的问题。正确的加载方式:
# custom_train.py
import torch
from yolox.models import YOLOX, YOLOPAFPN, YOLOXHead
def load_pretrained_weights(model, pretrained_path, num_classes):
"""
智能加载预训练权重,处理类别数不同的情况
"""
print(f"加载预训练权重: {pretrained_path}")
pretrained_dict = torch.load(pretrained_path, map_location="cpu")
if "model" in pretrained_dict:
pretrained_dict = pretrained_dict["model"]
model_dict = model.state_dict()
# 1. 过滤掉不匹配的键
matched_dict = {}
for k, v in pretrained_dict.items():
if k in model_dict:
# 检查维度是否匹配
if v.shape == model_dict[k].shape:
matched_dict[k] = v
else:
print(f"跳过 {k}: 维度不匹配 {v.shape} -> {model_dict[k].shape}")
# 对于分类头,如果是类别数不同,可以部分初始化
if "cls_preds" in k and len(v.shape) == 4:
# 保留共享的特征,只替换最后一层
min_channels = min(v.shape[1], model_dict[k].shape[1])
matched_weight = torch.zeros_like(model_dict[k])
matched_weight[:, :min_channels, :, :] = v[:, :min_channels, :, :]
matched_dict[k] = matched_weight
else:
print(f"跳过 {k}: 在模型中不存在")
# 2. 更新模型参数
model_dict.update(matched_dict)
model.load_state_dict(model_dict)
print(f"成功加载 {len(matched_dict)}/{len(model_dict)} 个参数")
return model
# 在训练脚本中使用
model = YOLOX(...)
model = load_pretrained_weights(model, "yolox_s.pth", num_classes=2)
4.2 训练参数调优指南
不同规模的数据集需要不同的训练策略。以下是我总结的经验配置:
| 数据集规模 | batch_size | 初始学习率 | warmup_epochs | max_epoch | 是否冻结骨干网络 |
|---|---|---|---|---|---|
| < 500张 | 4-8 | 0.0005 | 5 | 100-150 | 是(前50轮) |
| 500-2000张 | 8-16 | 0.001 | 5 | 150-200 | 是(前30轮) |
| 2000-5000张 | 16-32 | 0.01 | 5 | 200-300 | 否 |
| > 5000张 | 32-64 | 0.01 | 5 | 300+ | 否 |
对应的训练命令示例:
# 小数据集训练(冻结骨干网络)
python tools/train.py \
-f exps/example/yolox_voc/yolox_voc_s.py \
-d 1 -b 8 \
--fp16 \
-c yolox_s.pth \
--freeze_backbone 50 \ # 冻结前50轮
--min_lr_ratio 0.01 \ # 最小学习率为初始的1%
--warmup_epochs 5
# 大数据集训练(完整训练)
python tools/train.py \
-f exps/example/yolox_voc/yolox_voc_s.py \
-d 2 -b 32 \ # 使用2张GPU
--fp16 \
-c yolox_s.pth \
--min_lr_ratio 0.001 \
--warmup_epochs 5 \
--no_aug_epochs 15 \ # 最后15轮关闭数据增强
--ema # 启用指数移动平均
4.3 训练监控与调试技巧
训练过程中要实时监控,及时发现问题:
实时监控GPU使用情况:
# 在另一个终端运行
watch -n 1 nvidia-smi
# 或者使用更详细的监控
gpustat -i 1 # 每秒刷新一次
训练日志分析脚本:
# monitor_training.py
import re
import matplotlib.pyplot as plt
from pathlib import Path
def parse_training_log(log_file):
"""解析训练日志,提取关键指标"""
epochs = []
losses = []
lrs = []
with open(log_file, 'r') as f:
for line in f:
# 匹配训练损失
if 'Train loss:' in line:
match = re.search(r'Train loss: (\d+\.\d+)', line)
if match:
losses.append(float(match.group(1)))
# 匹配学习率
elif 'lr:' in line:
match = re.search(r'lr: (\d+\.\d+e?[-+]?\d*)', line)
if match:
lrs.append(float(match.group(1)))
# 匹配epoch
elif 'epoch:' in line and 'iter' in line:
match = re.search(r'epoch: (\d+)', line)
if match:
epochs.append(int(match.group(1)))
return epochs, losses, lrs
def plot_training_curves(log_dir):
"""绘制训练曲线"""
log_files = list(Path(log_dir).glob("*.log"))
fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(12, 8))
for log_file in log_files:
epochs, losses, lrs = parse_training_log(log_file)
if epochs and losses:
ax1.plot(epochs[:len(losses)], losses, label=log_file.stem)
ax1.set_ylabel('Loss')
ax1.set_xlabel('Epoch')
ax1.legend()
ax1.grid(True)
if epochs and lrs:
ax2.plot(epochs[:len(lrs)], lrs, label=log_file.stem)
ax2.set_ylabel('Learning Rate')
ax2.set_xlabel('Epoch')
ax2.legend()
ax2.grid(True)
plt.tight_layout()
plt.savefig('training_curves.png', dpi=150)
plt.show()
# 使用示例
plot_training_curves("./YOLOX_outputs/yolox_voc_s/")
常见训练问题与解决方案:
-
Loss为NaN或突然变大
# 在train.py中添加梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=10.0) # 或者降低学习率 self.basic_lr_per_img = 0.001 / 64.0 # 原来是0.01 -
GPU内存不足
# 减小batch_size -b 4 # 原来是8或16 # 减小输入尺寸 self.input_size = (416, 416) # 原来是640 # 使用梯度累积 self.accumulate = 4 # 每4个batch更新一次权重 -
训练速度慢
# 增加num_workers(但不要超过CPU核心数) self.data_num_workers = 8 # 使用更快的DataLoader self.persistent_workers = True # 启用pin_memory(如果内存足够) self.pin_memory = True
5. 模型评估与部署:从验证到生产
训练完成后,正确的评估和部署同样重要。很多人在这里遇到模型加载失败、精度下降的问题。
5.1 多指标评估与模型选择
YOLOX默认只输出mAP,但实际项目中可能需要更多指标:
# custom_eval.py
from yolox.evaluators import VOCEvaluator
import numpy as np
from collections import defaultdict
class EnhancedVOCEvaluator(VOCEvaluator):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
self.per_class_metrics = defaultdict(list)
def evaluate(self, model, distributed=False, half=False):
"""增强的评估方法,输出每类指标"""
results = super().evaluate(model, distributed, half)
# 解析每类AP
if hasattr(self, 'dataloader'):
# 这里需要根据实际数据集获取类别信息
class_names = self.dataloader.dataset.classes
print("\n" + "="*50)
print("每类性能分析:")
print("="*50)
# 实际项目中需要从结果中提取每类AP
# 这里只是示例结构
for i, class_name in enumerate(class_names):
# 假设有每类的AP数据
ap = 0.85 # 示例值
print(f"{class_name:20s} AP@0.5: {ap:.3f}")
# 记录用于后续分析
self.per_class_metrics[class_name].append(ap)
return results
# 使用示例
evaluator = EnhancedVOCEvaluator(
dataloader=val_loader,
img_size=test_size,
confthre=0.01,
nmsthre=0.65,
num_classes=num_classes,
testdev=False,
)
metrics = evaluator.evaluate(model)
5.2 模型导出与优化
YOLOX支持多种导出格式,但每个都有注意事项:
导出ONNX:
python tools/export_onnx.py \
--output-name yolox_s.onnx \
-f exps/example/yolox_voc/yolox_voc_s.py \
-c YOLOX_outputs/yolox_voc_s/best_ckpt.pth \
--dynamic # 支持动态batch size
注意:ONNX导出时可能会遇到不支持的算子。如果出错,尝试:
- 更新torch和onnx版本
- 添加
--opset 12参数指定opset版本- 简化模型结构,移除自定义算子
TensorRT优化:
# trt_optimize.py
import tensorrt as trt
import onnx
import os
def build_engine(onnx_file, engine_file, max_batch_size=1, fp16_mode=True):
"""构建TensorRT引擎"""
TRT_LOGGER = trt.Logger(trt.Logger.WARNING)
builder = trt.Builder(TRT_LOGGER)
network = builder.create_network(
1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)
)
parser = trt.OnnxParser(network, TRT_LOGGER)
# 解析ONNX模型
with open(onnx_file, 'rb') as model:
if not parser.parse(model.read()):
print("解析ONNX失败:")
for error in range(parser.num_errors):
print(parser.get_error(error))
return None
# 配置优化参数
config = builder.create_builder_config()
config.max_workspace_size = 1 << 30 # 1GB
if fp16_mode and builder.platform_has_fast_fp16:
config.set_flag(trt.BuilderFlag.FP16)
# 设置动态shape
profile = builder.create_optimization_profile()
input_tensor = network.get_input(0)
# 最小、最优、最大shape
profile.set_shape(
input_tensor.name,
(1, 3, 416, 416), # 最小
(max_batch_size, 3, 640, 640), # 最优
(max_batch_size, 3, 960, 960) # 最大
)
config.add_optimization_profile(profile)
# 构建引擎
print("开始构建TensorRT引擎...")
engine = builder.build_engine(network, config)
if engine is None:
print("构建引擎失败")
return None
# 保存引擎
with open(engine_file, "wb") as f:
f.write(engine.serialize())
print(f"引擎已保存到: {engine_file}")
return engine
# 使用示例
build_engine("yolox_s.onnx", "yolox_s.engine", max_batch_size=4, fp16_mode=True)
5.3 部署时的性能优化
在生产环境中部署时,还需要考虑:
批处理优化:
# batch_inference.py
import torch
import time
from yolox.utils import postprocess
class BatchInference:
def __init__(self, model, batch_size=4, img_size=640):
self.model = model
self.batch_size = batch_size
self.img_size = img_size
self.buffer = []
def preprocess_batch(self, images):
"""批量预处理"""
batch_tensors = []
for img in images:
# 这里添加你的预处理逻辑
tensor = self.preprocess_single(img)
batch_tensors.append(tensor)
return torch.stack(batch_tensors)
def inference(self, images):
"""批量推理"""
if len(images) == 0:
return []
# 累积到batch_size再推理
self.buffer.extend(images)
if len(self.buffer) >= self.batch_size:
batch = self.buffer[:self.batch_size]
self.buffer = self.buffer[self.batch_size:]
# 批量处理
batch_tensor = self.preprocess_batch(batch)
with torch.no_grad():
if torch.cuda.is_available():
batch_tensor = batch_tensor.cuda()
outputs = self.model(batch_tensor)
results = postprocess(
outputs,
num_classes=self.model.num_classes,
conf_thre=0.3,
nms_thre=0.65
)
return results
return None # 等待更多数据
def flush(self):
"""处理剩余数据"""
if self.buffer:
return self.inference(self.buffer)
return []
内存管理策略:
# memory_manager.py
import gc
import torch
class InferenceMemoryManager:
def __init__(self, model, max_memory_mb=1024):
self.model = model
self.max_memory = max_memory_mb * 1024 * 1024 # 转换为字节
def adaptive_batch_size(self, input_shape):
"""根据输入尺寸自适应调整batch_size"""
# 估算单张图片的内存占用
single_image_memory = self.estimate_memory(input_shape)
# 计算最大batch_size
max_batch = int(self.max_memory / single_image_memory)
# 留出安全余量
safe_batch = max(1, max_batch - 2)
print(f"输入尺寸: {input_shape}, 单张内存: {single_image_memory/1024/1024:.1f}MB")
print(f"建议batch_size: {safe_batch}")
return safe_batch
def estimate_memory(self, input_shape):
"""估算内存占用"""
# 模型参数内存
param_memory = sum(p.numel() * p.element_size()
for p in self.model.parameters())
# 激活内存(近似估算)
batch, channels, height, width = input_shape
# 这里简化计算,实际需要根据模型结构估算
activation_memory = batch * channels * height * width * 4 # 假设float32
return param_memory + activation_memory
def cleanup(self):
"""清理内存"""
torch.cuda.empty_cache()
gc.collect()
6. 实战案例:工业缺陷检测全流程
让我分享一个真实的工业缺陷检测项目经验。客户需要检测PCB板上的焊接缺陷,数据集只有800张标注图片,但要求检测精度达到95%以上。
6.1 数据层面的挑战与解决方案
挑战1:样本不均衡
- 正常样本:650张
- 缺陷样本:150张(其中某种罕见缺陷只有20张)
解决方案:
# 使用加权采样
from torch.utils.data import WeightedRandomSampler
def create_balanced_sampler(dataset):
"""创建平衡采样器"""
# 计算每个类别的样本数
class_counts = count_classes(dataset)
# 计算每个样本的权重
weights = []
for _, target in dataset:
class_id = target[0] # 假设第一个是类别
weight = 1.0 / class_counts[class_id]
weights.append(weight)
sampler = WeightedRandomSampler(
weights,
num_samples=len(weights),
replacement=True
)
return sampler
# 在数据加载器中使用
train_loader = DataLoader(
dataset=train_dataset,
batch_size=8,
sampler=balanced_sampler,
num_workers=4,
pin_memory=True
)
挑战2:小目标检测 缺陷尺寸很小,有的只占图像的1%。
解决方案:
# 修改模型配置
class Exp(MyExp):
def __init__(self):
super(Exp, self).__init__()
# 使用更密集的锚点
self.use_l1 = False
self.use_aux_loss = True # 启用辅助损失
# 修改检测头,增加小目标检测能力
self.depthwise = False
self.act = 'silu'
# 数据增强增强小目标
self.mosaic_prob = 1.0
self.mixup_prob = 0.5
self.hsv_prob = 1.0
self.flip_prob = 0.5
# 更小的anchor尺寸
self.multiscale_range = 3 # 多尺度训练范围
6.2 训练策略调整
针对小数据集,我采用了以下策略:
- 渐进式解冻:
def progressive_unfreeze(model, epoch, total_epochs):
"""渐进式解冻策略"""
# 第1-30轮:只训练检测头
if epoch < 30:
for name, param in model.named_parameters():
if 'backbone' in name:
param.requires_grad = False
else:
param.requires_grad = True
# 第31-100轮:解冻最后两个阶段
elif epoch < 100:
for name, param in model.named_parameters():
if 'backbone.stage4' in name or 'backbone.stage3' in name:
param.requires_grad = True
elif 'backbone' in name:
param.requires_grad = False
else:
param.requires_grad = True
# 第100轮后:解冻全部
else:
for param in model.parameters():
param.requires_grad = True
- 学习率预热与衰减:
def adjust_learning_rate(optimizer, epoch, warmup_epochs, total_epochs):
"""自定义学习率调度"""
if epoch < warmup_epochs:
# 线性预热
lr = base_lr * (epoch + 1) / warmup_epochs
elif epoch < total_epochs * 0.8:
# 余弦衰减
progress = (epoch - warmup_epochs) / (total_epochs * 0.8 - warmup_epochs)
lr = base_lr * 0.5 * (1 + math.cos(math.pi * progress))
else:
# 最后阶段保持小学习率
lr = base_lr * 0.01
for param_group in optimizer.param_groups:
param_group['lr'] = lr
6.3 结果与优化
经过上述调整,最终在测试集上达到了96.3%的mAP。关键改进包括:
- 使用Focal Loss处理类别不均衡
- 添加注意力机制增强小目标检测
- 采用Test Time Augmentation (TTA) 提升推理精度
- 模型蒸馏压缩模型大小
最终的推理脚本:
# inference_with_tta.py
import torch
import torchvision.transforms as T
from yolox.utils import postprocess
class TTAInference:
"""测试时增强推理"""
def __init__(self, model, num_classes, scales=[0.8, 1.0, 1.2]):
self.model = model
self.num_classes = num_classes
self.scales = scales
def tta_transform(self, image):
"""生成多个增强版本"""
transforms = []
original_size = image.size
# 不同尺度的缩放
for scale in self.scales:
new_size = (int(original_size[0] * scale),
int(original_size[1] * scale))
transforms.append(T.Resize(new_size))
# 水平翻转
transforms.append(T.Compose([
T.Resize(original_size),
T.RandomHorizontalFlip(p=1.0)
]))
return transforms
def inference(self, image, conf_thre=0.3, nms_thre=0.65):
"""TTA推理"""
all_predictions = []
for transform in self.tta_transform(image):
transformed_img = transform(image)
tensor_img = self.preprocess(transformed_img)
with torch.no_grad():
if torch.cuda.is_available():
tensor_img = tensor_img.cuda()
output = self.model(tensor_img)
predictions = postprocess(
output,
self.num_classes,
conf_thre,
nms_thre
)
# 将预测框转换回原始尺寸
if predictions[0] is not None:
predictions = self.inverse_transform(
predictions, transform, original_size
)
all_predictions.append(predictions)
# 融合所有预测结果
if all_predictions:
final_predictions = self.merge_predictions(all_predictions)
return final_predictions
return None
这个项目让我深刻体会到,YOLOX虽然强大,但要发挥其全部潜力,需要根据具体任务精心调整。每个数据集都有其独特性,没有一刀切的解决方案。关键是要理解模型的工作原理,然后针对性地解决数据、训练、部署中的实际问题。
&spm=1001.2101.3001.5002&articleId=153724412&d=1&t=3&u=e0ba2195f0494026993e75484ded3cb6)
75

被折叠的 条评论
为什么被折叠?



