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

Yolo-NAS自定义数据训练时出现张量维度不匹配RuntimeError

Yolo-NAS训练RuntimeError排查与解决

错误原因

错误提示里的tensor a (80)是COCO预训练模型默认的80个类别维度,tensor b (3)是你的自定义数据集类别数,两者维度不匹配。核心问题是加载预训练模型时未修改输出头的类别数,导致模型仍输出80类的预测结果,但损失函数和数据集是3类,计算损失时出现维度冲突。

解决方法

1. 修正模型加载代码

加载Yolo-NAS模型时必须指定num_classes参数,匹配自定义数据集的类别数量,替换原有模型加载代码:

import torch
from super_gradients.training import models

DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
MODEL_ARCH = 'yolo_nas_m'

# 关键:添加num_classes参数,指定自定义类别数
model = models.get(MODEL_ARCH, pretrained_weights="coco", num_classes=len(CLASSES)).to(DEVICE)

这样模型会自动调整输出层维度,适配你的3类数据集。

2. 验证数据集标注正确性

确认自定义数据集的标注文件满足:

  • 类别索引从0开始,最大索引不超过len(CLASSES)-1(即2)
  • 标注格式符合YOLO标准(每行内容为class_index x_center y_center width height)

补充说明

你的训练参数中已经正确设置了损失函数和指标的num_classes/num_cls,但模型本身的输出维度未修改,这是触发错误的核心原因。修改模型加载代码后,即可解决维度不匹配的RuntimeError。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 15:14:56