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 ...
问题原因
- 保存格式不兼容:YOLOv8的
attempt_load_one_weight函数要求加载的检查点必须包含model或ema顶层键,但你用torch.save(model.state_dict())保存的只是单纯的权重参数字典,缺少这些关键结构,导致加载时找不到ckpt["model"]。 - 剪枝匹配逻辑有风险:当前通过
"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
相关产品推荐
相关产品推荐

