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

YOLOv8检测模型剪枝后训练/预测报错问题排查

YOLOv8剪枝后模型加载报错问题排查

报错回溯

Traceback (most recent call last):
  File "C:\ProgramData\anaconda3\lib\runpy.py", line 196, in _run_module_as_main
    return _run_code(code, main_globals, None,
  File "C:\ProgramData\anaconda3\lib\runpy.py", line 86, in _run_code
    exec(code, run_globals)
  File "C:\ProgramData\anaconda3\Scripts\yolo.exe\__main__.py", line 7, in <module>
  File "C:\Users\user\AppData\Roaming\Python\Python310\site-packages\ultralytics\cfg\__init__.py", line 555, in entrypoint
    model = YOLO(model, task=task)
  File "C:\Users\user\AppData\Roaming\Python\Python310\site-packages\ultralytics\models\yolo\model.py", line 23, in __init__
    super().__init__(model=model, task=task, verbose=verbose)
  File "C:\Users\user\AppData\Roaming\Python\Python310\site-packages\ultralytics\engine\model.py", line 151, in __init__
    self._load(model, task=task)
  File "C:\Users\user\AppData\Roaming\Python\Python310\site-packages\ultralytics\engine\model.py", line 240, in _load
    self.model, self.ckpt = attempt_load_one_weight(weights)
  File "C:\Users\user\AppData\Roaming\Python\Python310\site-packages\ultralytics\nn\tasks.py", line 792, in attempt_load_one_weight
    model = (ckpt.get("ema") or ckpt["model"]).to(device).float()  # FP32 model

剪枝代码

import torch
from torch.nn.utils import prune
from ultralytics import YOLO

# Load your model
model = YOLO('best.pt')
ratio=0.5

for name, module in model.named_modules():
    if "cv1.conv" in name:  # Check if the layer name contains "cv1"
        print(f"Pruning layer: {name}")
        prune.l1_unstructured(module, name='weight', amount=ratio)  # Prune the 'conv' submodule
        prune.remove(module, 'weight')  # Optional: Apply pruning permanently

# Save the pruned model
torch.save(model.state_dict(), 'pruned_model.pt')

模型部分层结构

...
model.model.18
model.model.18.cv1
model.model.18.cv1.conv
model.model.18.cv1.bn
model.model.18.cv2
model.model.18.cv2.conv
model.model.18.cv2.bn
...
model.model.19.conv
model.model.19.bn
model.model.20
model.model.21
...

问题原因

  1. 保存格式不兼容:YOLOv8的attempt_load_one_weight函数要求加载的检查点必须包含model或ema顶层键,但你用torch.save(model.state_dict())保存的只是单纯的权重参数字典,缺少这些关键结构,导致加载时找不到ckpt["model"]。
  2. 剪枝匹配逻辑有风险:当前通过"cv1.conv" in name匹配目标层,若模型中存在其他名称包含该字符串的模块,会出现误剪情况;且剪枝后未遵循YOLOv8的模型保存规范,进一步加剧加载失败问题。

修正方案

  • 改用YOLOv8原生保存方法:替换torch.save(model.state_dict(), 'pruned_model.pt')为model.save('pruned_model.pt'),该方法会保存包含model、ema等必要键的完整检查点,完全适配YOLOv8的加载逻辑。
  • 优化剪枝匹配逻辑:结合模块类型和名称双重判断,精准定位目标层,避免误剪:
    for name, module in model.named_modules():
        # 仅对cv1下的卷积层执行剪枝
        if isinstance(module, torch.nn.Conv2d) and ".cv1." in name:
            print(f"Pruning layer: {name}")
            prune.l1_unstructured(module, name='weight', amount=ratio)
            prune.remove(module, 'weight')
    
  • 验证加载流程:剪枝并保存后,直接用YOLO('pruned_model.pt')加载模型,测试预测或训练是否正常运行。

内容的提问来源于stack exchange,提问作者Mary H

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 13:50:49