一、训练具体教会AI什么

可以把一个AI模型(无论是YOLO还是语言模型)想象成一个初始只有基础结构、但一片空白的“数字大脑”。训练的本质,就是通过反复给它看“数据教材”,来调整它内部数以百万甚至亿计的“神经连接”的强弱,从而让它形成对特定任务的“条件反射”或“本能反应”。

二、环境配置

2.1安装虚拟环境

在终端进入训练的文件夹中安装虚拟环境,我电脑的python版本是3.10.9,这里也安装这个版本

conda create -n yolo_train python=3.10.9 -y

可以看到已经创建成功。

2.2添加python解释器

把项目拖进pycharm中,然后在右下角点击【添加新的解释器】或者直接在【解释器设置中设置】,设置完成后点击应用。

右下角切换成功即可。

2.3激活前面创建的虚拟环境

回到命令行中,先激活虚拟环境。

conda activate yolo_train

2.4在虚拟环境中安装所有必需的库

# 安装PyTorch深度学习框架(根据你的电脑无NVIDIA GPU的情况,安装CPU版本)
conda install pytorch torchvision torchaudio cpuonly -c pytorch -y
# 安装Ultralytics YOLO库,这是一个非常易用的目标检测框架
pip install ultralytics
# 安装常用的数据科学工具包
pip install opencv-python pandas matplotlib jupyter

2.5准备项目目录结构与数据


your_project/
├── datasets/
│   └── my_dataset/       # 数据集
│       ├── images/
│       │   ├── train/    # 存放训练图片
│       │   └── val/      # 存放验证图片
│       └── labels/
│           ├── train/    # 存放对应的训练标签文件 (.txt)
│           └── val/      # 存放对应的验证标签文件
├── data.yaml             # 数据集配置文件(核心!)
└── train.py             # 训练脚本

2.6编写训练脚本 train.py

# train.py
from ultralytics import YOLO
import os


def main():
    print("=" * 50)
    print("开始动物检测模型训练")
    print("=" * 50)

    # 检查配置文件是否存在
    config_file = 'data.yaml'
    if not os.path.exists(config_file):
        print(f"❌ 错误:找不到配置文件 {config_file}")
        print("请确保data.yaml文件与train.py在同一目录")
        return

    print(f"✅ 找到配置文件: {config_file}")

    # 检查数据集结构
    required_dirs = [
        'datasets/my_dataset/images/train',
        'datasets/my_dataset/images/val',
        'datasets/my_dataset/labels/train',
        'datasets/my_dataset/labels/val'
    ]

    print("\n🔍 检查数据集结构...")
    all_exists = True
    for dir_path in required_dirs:
        if os.path.exists(dir_path):
            file_count = len([f for f in os.listdir(dir_path) if f.endswith(('.jpg', '.jpeg', '.png', '.txt'))])
            print(f"  ✅ {dir_path}: 找到 {file_count} 个文件")
        else:
            print(f"  ❌ {dir_path}: 目录不存在")
            all_exists = False

    if not all_exists:
        print("\n⚠️  请确保数据集结构正确后再继续")
        return

    print("\n✅ 数据集结构检查完成")

    # 加载预训练模型(会自动下载)
    print("\n📥 加载YOLOv8n预训练模型...")
    try:
        model = YOLO('yolov8n.pt')  # 使用最小的nano版本,适合CPU训练
        print("✅ 模型加载成功")
    except Exception as e:
        print(f"❌ 模型加载失败: {e}")
        print("请检查网络连接,或手动下载模型文件")
        return

    # 训练配置(针对CPU配置优化)
    print("\n⚙️  配置训练参数...")
    print("  设备: CPU (检测到无NVIDIA GPU)")
    print("  图片尺寸: 640x640")
    print("  批次大小: 1")
    print("  训练轮数: 50")
    print("  类别数量: 5 (狗、猪、猫、鸟、蛇)")

    # 开始训练!
    print("\n🚀 开始训练模型...")
    print("训练过程可能需要几分钟到几小时,请耐心等待")
    print("可以观察损失值(loss)下降趋势判断训练是否正常")
    print("-" * 50)

    try:
        results = model.train(
            data=config_file,  # 数据集配置文件
            epochs=50,  # 训练轮数
            imgsz=640,  # 输入图片尺寸
            batch=1,  # 批次大小(CPU训练建议4或更小)
            device='cpu',  # 使用CPU训练
            workers=0,  # Windows系统设为0避免问题
            patience=15,  # 早停耐心值
            save=True,  # 保存训练结果
            save_period=10,  # 每10轮保存一次检查点
            name='my_dataset_v1',  # 实验名称
            project='runs',  # 结果保存目录
            verbose=True,  # 显示详细输出
            val=True,  # 启用验证
            exist_ok=True  # 允许覆盖现有结果
        )

        print("\n" + "=" * 50)
        print("🎉 训练完成!")
        print("=" * 50)

        # 显示训练结果保存位置
        print("\n📁 训练结果保存在以下目录:")
        print("  runs/detect/my_dataset_v1/")
        print("\n📊 重要文件:")
        print("  - weights/best.pt: 最佳模型权重")
        print("  - weights/last.pt: 最终模型权重")
        print("  - results.png: 训练曲线图")
        print("  - args.yaml: 训练参数备份")

        # 验证训练效果
        print("\n🔍 验证模型性能...")
        print("运行以下命令测试模型:")
        print("python test_model.py")

    except KeyboardInterrupt:
        print("\n⚠️  训练被用户中断")
        print("部分结果已保存,可以使用 last.pt 继续训练")
    except Exception as e:
        print(f"\n❌ 训练过程中出现错误: {e}")
        print("请检查错误信息并修正后重试")


if __name__ == '__main__':
    main()

2.7创建测试脚本 test_model.py

# test_model.py
from ultralytics import YOLO
import cv2
import os

def test_trained_model():
    print("测试训练好的模型...")
    
    # 加载训练得到的最佳模型
    model_path = 'runs/detect/animal_detection_v1/weights/best.pt'
    
    if not os.path.exists(model_path):
        print(f"❌ 找不到模型文件: {model_path}")
        print("请先完成训练,或检查路径是否正确")
        return
    
    model = YOLO(model_path)
    print(f"✅ 加载模型: {model_path}")
    
    # 使用验证集中的一张图片进行测试
    test_image_path = 'datasets/animal_detection/images/val/'
    
    # 自动查找验证集中的第一张图片
    if os.path.exists(test_image_path):
        image_files = [f for f in os.listdir(test_image_path) 
                      if f.lower().endswith(('.jpg', '.jpeg', '.png'))]
        
        if image_files:
            test_image = os.path.join(test_image_path, image_files[0])
            print(f"🔍 测试图片: {test_image}")
            
            # 进行预测
            results = model.predict(
                source=test_image,
                save=True,           # 保存结果
                save_txt=True,       # 保存标签
                conf=0.25,           # 置信度阈值
                save_conf=True       # 保存置信度
            )
            
            # 显示结果
            for result in results:
                result.show()  # 显示预测结果
                
                print("\n📋 检测结果:")
                if result.boxes is not None:
                    for box in result.boxes:
                        class_id = int(box.cls[0])
                        confidence = float(box.conf[0])
                        class_name = model.names[class_id]
                        print(f"  - {class_name}: {confidence:.2%}")
                else:
                    print("  未检测到目标")
        else:
            print("❌ 验证集中找不到图片文件")
    else:
        print(f"❌ 验证集路径不存在: {test_image_path}")

if __name__ == '__main__':
    test_trained_model()

2.8编写data.yaml

# data.yaml
# 数据集路径(相对于此yaml文件的位置)
path: ./datasets/animal_detection
train: images/train
val: images/val

# 类别数量 (nc = number of classes)
nc: 5

# 类别名称列表 (必须与标注时使用的类别名称完全一致,且顺序固定)
names: ['dog','pig', 'cat', 'bird', 'snake']
# 注意:此处的索引顺序就是class_id,即 dog=0, pig=1, cat=2, bird=3, snake=4

2.9运行训练脚本train.py

cd /path/to/your/project
python train.py

这里运行的50张图片,电脑配置不行,跑挺久😵‍💫

代码输出的结果:

要注意检查ai帮忙写的代码路径可能对不上,自己还要再检查修改

2.10运行test_model.py

python test_model.py

运行结果:运行完成后会自动弹出这个图片,弹出来的这个图片倒是准确,不过其他的准确性不是很高,还有没有识别出的🤣

2.11结果解读

第一次操作,生成的图片找AI帮忙解读情况。

1、weights/best.pt: 最佳模型权重

best.pt本质:它是一个 PyTorch 模型权重文件.pt.pth是PyTorch框架用于保存模型状态的标准文件扩展名。

内容:这个文件二进制形式存储了刚才训练的YOLO模型在训练结束后被认为在验证集上性能最佳(通常是mAP最高)时的全部可学习参数

来源:在运行 train.py脚本后,Ultralytics YOLO框架会在每轮训练后自动在验证集上评估模型。在整个训练过程中,表现最好的那个模型状态会被保存为 best.pt。同时,训练结束时的最后一个状态会被保存为 last.pt

**用处:模型部署与推理。**可以使用这个文件,让模型对新的、从未见过的图片进行预测。

**模型调优的起点:**如果后面想用更多数据继续训练这个模型(即“微调”),您可以加载 best.pt作为预训练权重开始训练,而不是从零开始(yolov8n.pt),这样可以更快收敛并可能获得更好性能。

2、weights/last.pt: 最终模型权重

定义last.pt是模型在设定的全部训练轮次(Epoch)结束后那一刻的完整状态快照

可以加载 last.pt并评估其性能,与 best.pt的评估结果进行对比。如果 last.pt的指标显著低于 best.pt,这是一个强烈的信号,表明模型在训练后期很可能出现了过拟合

3、results.png: 训练曲线图

核心结论是:训练是有效且正常的,模型确实在学习,但受限于较小的数据集,其泛化能力(处理新图片的能力)有提升空间。

📊 整体训练状态解读

指标类型 曲线趋势 解读 对您模型的评估
训练损失 (Train Losses) 所有指标(框/分类/特征损失)均快速下降后趋于平缓 模型在“练习题”(训练集)上越做越熟练,正在有效学习。 优秀。表明模型有能力从我提供的数据中学习规律。
验证损失 (Val Losses) 前期下降后,在中后期进入波动/平台期,未继续显著下降。 模型在“模拟考”(验证集)上表现提升乏力,可能已学到当前数据的所有规律。 ⚠️ 正常,但提示上限。这是小数据集的典型表现,模型“学无可学”了。
验证精度 (mAP) mAP50 达到 ~0.85 后在高位剧烈波动;mAP50-95 较低且在 ~0.35 波动。 模型能以高置信度检测出目标(高mAP50),但框的位置不够精准(低mAP50-95)。 ⚠️ 符合小数据预期。模型学会了“是什么”,但对“精确在哪”把握不足。

4、args.yaml: 训练参数备份

(1) 训练过程控制参数

参数 默认/设置的值 解释与作用
**epochs** 50 训练总轮数。模型将完整遍历训练数据集多少次。轮数太少模型学不充分,太多可能导致过拟合。可观察验证集损失曲线来确定最佳轮数。
**batch** 4 批次大小。一次迭代中输入模型的图片张数。受GPU/CPU内存限制。增大batch可能使训练更稳定,但内存占用线性增长。
**patience** 15 早停耐心值。如果验证集指标连续patience轮没有提升,则自动停止训练以防止过拟合。这是防止无效训练的重要“安全阀”。
**save** true 是否保存模型。会保存best.ptlast.pt
**save_period** -1 定期保存检查点的周期。-1表示不按周期保存,只保存最佳和最终模型。设为10则每10轮额外保存一个检查点。

(2)优化与正则化参数(调优重点)

参数 默认/您的值 解释与作用
**lr0** 0.01 初始学习率 (Initial Learning Rate)。控制模型参数每次更新的步长。最重要的超参数之一。太大可能导致训练震荡或不收敛,太小则训练缓慢。通常需要根据任务调整。
**lrf** 0.01 最终学习率因子 (Final Learning Rate Factor)。训练结束时,学习率将衰减为 lr0 * lrf。设置为0.01,意味着最终学习率是初始值的1%(0.0001)。这属于余弦退火等学习率调度策略的一部分,有助于模型后期精细调优。
**momentum** 0.937 动量。优化器(如SGD with momentum)参数,帮助加速收敛并冲出局部最优点。通常保持默认。
**weight_decay** 0.0005 权重衰减。一种L2正则化,对大的模型参数施加惩罚,是防止模型过拟合的核心技术之一。值越大,正则化越强。
**augment** false 是否启用数据增强建议在后续训练中设为**true**。它将随机对训练图片进行翻转、裁剪、色彩变化等,能有效提升模型泛化能力,是应对您当前小数据集过拟合风险的最实用技巧。

三、疑问

ps.虽然跟着视频操作加上AI指导,但还是很多疑问,如模型预测是怎么实现的?

先简单总结预测流程加载模型 -> 预处理输入 -> 执行前向传播-> 后处理输出 -> 呈现结果

(1)模型加载:实例化模型对象。在内存中创建一个可计算的对象,它具备了接收输入、进行计算、产生输出的完整能力。

(2)输入数据预处理:当传入一张新图片时,预测流程不会直接处理原始图片。系统会先对其进行一系列标准化操作,以匹配模型训练时所期望的输入格式如调整图片尺寸归一化、张量转换

(3)前向传播:模型预测的核心计算过程

(4)后处理:模型直接输出的原始张量是机器友好的,但人难以理解。因此需要进行后处理,将其转化为直观的结果。

  • 非极大值抑制:如果模型对同一个目标预测了多个重叠的框,此算法会只保留其中置信度最高的一个,消除冗余框。
  • 置信度过滤:舍弃那些模型自己都“没把握”(置信度低于阈值,如0.5)的预测框。
  • 结果格式化:将最终保留的预测框信息(边界框的坐标、类别标签、置信度)转换为结构化的数据(如列表、字典)或可视化结果(在图片上绘制出带有标签和置信度的矩形框)。
Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐