使用PyTorch框架和YOLOv5模型训练矿场矿车坐人数据集 来识别乘坐的煤矿矿车、正常情况下的矿车人员乘坐情况 构建矿场矿车坐人检测深度学习模型
文章目录
- 使用PyTorch框架和YOLOv5模型训练矿场矿车坐人数据集 来识别乘坐的煤矿矿车、正常情况下的矿车人员乘坐情况 构建矿场矿车坐人检测深度学习模型
- 表格表示
- 1. 环境设置
- 2. 数据准备
- 创建自定义Dataset类
- 3. 数据增强和转换
- 4. 模型定义
- 5. 训练流程
- 6. 模型保存与加载
- 7. 应用到实际场景中
- 注意事项
以下文字及代码仅供参考。
矿场矿车坐人数据集 3900张 矿车坐人 voc yolo 标注
三类:有人员乘坐的矿车、正常情况下的矿车(没有人员乘坐),以及异常情况下的矿车。共8475个边界框标注,具体分布如下表所示。
表格表示
| 类别 | 图片数量 | 边界框数量 |
|---|---|---|
| 有人员乘坐 | 344 | 404 |
| 正常情况 | 3335 | 7739 |
| 异常情况 | 272 | 332 |
| 总计 | 3963 | 8475 |
备注提醒:其中,“有人员乘坐”指的是矿车上有人的情况;“正常情况”指的是矿车按预期使用但无人乘坐的情况;而“异常情况”可能指的是除上述两种情况之外的特殊情况,例如设备故障或其他非标准操作状态。这里的数据总计与提供的数字略有出入,可能是由于四舍五入或额外的数据点造成的,因此请根据实际情况进行调整。
labe|| pic_num| boxnum
With People:(344,404)
Norma | :(3335,7739)
Abnormal: (272, 332)
total :(3963,8475)
1
构建一个矿场矿车坐人检测系统涉及多个步骤,包括数据预处理、模型选择与训练、评估和优化。我们将使用PyTorch框架和YOLOv5模型作为示例。以下是详细的步骤和代码:
1. 环境设置
确保安装了必要的库:
pipinstalltorch torchvision opencv-python numpy pipinstallyolov52. 数据准备
假设你的数据集已经按照VOC格式标注好,并且包含以下目录结构:
data/ ├── images/ │ ├── train/ │ └── val/ └── labels/ ├── train/ └── val/创建自定义Dataset类
importosimportcv2importtorchfromtorch.utils.dataimportDatasetclassMineCartDataset(Dataset):def__init__(self,root,split='train',transform=None):self.root=root self.split=split self.transform=transform self.images_dir=os.path.join(root,'images',split)self.labels_dir=os.path.join(root,'labels',split)self.image_paths=sorted(os.listdir(self.images_dir))self.label_paths=[f.replace('.jpg','.txt')forfinself.image_paths]def__len__(self):returnlen(self.image_paths)def__getitem__(self,idx):img_path=os.path.join(self.images_dir,self.image_paths[idx])label_path=os.path.join(self.labels_dir,self.label_paths[idx])image=cv2.imread(img_path)image=cv2.cvtColor(image,cv2.COLOR_BGR2RGB)withopen(label_path,'r')asf:lines=f.readlines()boxes=[]forlineinlines:parts=line.strip().split()x_center,y_center,width,height=map(float,parts[1:])xmin=int((x_center-width/2)*image.shape[1])ymin=int((y_center-height/2)*image.shape[0])xmax=int((x_center+width/2)*image.shape[1])ymax=int((y_center+height/2)*image.shape[0])boxes.append([xmin,ymin,xmax,ymax])target={'boxes':torch.tensor(boxes,dtype=torch.float32),'labels':torch.ones(len(boxes),dtype=torch.int64)# 假设只有一个类别}ifself.transform:image,target=self.transform(image,target)returnimage,target3. 数据增强和转换
定义一些基本的数据增强和转换操作:
importalbumentationsasAfromalbumentations.pytorchimportToTensorV2defget_transform(train):iftrain:returnA.Compose([A.HorizontalFlip(p=0.5),A.RandomBrightnessContrast(p=0.2),A.Normalize(),ToTensorV2()])else:returnA.Compose([A.Normalize(),ToTensorV2()])4. 模型定义
使用YOLOv5模型进行训练:
importtorchimporttorch.nnasnnfromyolov5.models.commonimportDetectMultiBackendfromyolov5.utils.generalimportnon_max_suppressiondefget_model(num_classes):model=DetectMultiBackend('yolov5s.pt')model.nc=num_classes# number of classesmodel.class_names=['person']# class namesreturnmodel5. 训练流程
定义训练循环,并使用train_one_epoch和evaluate函数来训练和评估模型:
fromtorch.utils.dataimportDataLoaderfromyolov5.utils.datasetsimportLoadImagesfromyolov5.utils.torch_utilsimportselect_devicedeftrain(model,dataloader,optimizer,device):model.train()forimages,targetsindataloader:images=images.to(device)targets=[{k:v.to(device)fork,vint.items()}fortintargets]loss_dict=model(images,targets)losses=sum(lossforlossinloss_dict.values())optimizer.zero_grad()losses.backward()optimizer.step()defevaluate(model,dataloader,device):model.eval()withtorch.no_grad():forimages,targetsindataloader:images=images.to(device)targets=[{k:v.to(device)fork,vint.items()}fortintargets]outputs=model(images)predictions=non_max_suppression(outputs,conf_thres=0.5,iou_thres=0.5)# 这里可以添加评估指标的计算pass# 加载数据集dataset_train=MineCartDataset('path/to/data',split='train',transform=get_transform(train=True))dataloader_train=DataLoader(dataset_train,batch_size=4,shuffle=True,num_workers=4)dataset_val=MineCartDataset('path/to/data',split='val',transform=get_transform(train=False))dataloader_val=DataLoader(dataset_val,batch_size=4,shuffle=False,num_workers=4)# 初始化模型device=select_device('')model=get_model(num_classes=2)# 背景+无人机model.to(device)# 构建优化器optimizer=torch.optim.SGD(model.parameters(),lr=0.005,momentum=0.9,weight_decay=0.0005)# 开始训练num_epochs=10forepochinrange(num_epochs):train(model,dataloader_train,optimizer,device)evaluate(model,dataloader_val,device)6. 模型保存与加载
在训练结束后,可以保存模型以便后续使用:
torch.save(model.state_dict(),'mine_cart_detection_model.pth')# 加载模型model.load_state_dict(torch.load('mine_cart_detection_model.pth'))7. 应用到实际场景中
加载模型并进行推理:
defpreprocess_image(image_path):image=cv2.imread(image_path)image=cv2.cvtColor(image,cv2.COLOR_BGR2RGB)returnimagedefdetect_people(model,image_tensor,threshold=0.5):withtorch.no_grad():outputs=model(image_tensor.unsqueeze(0))predictions=non_max_suppression(outputs,conf_thres=threshold,iou_thres=0.5)returnpredictionsdefvisualize_predictions(image,predictions):forpredinpredictions:boxes=pred[:,:4].cpu().numpy()scores=pred[:,4].cpu().numpy()forbox,scoreinzip(boxes,scores):xmin,ymin,xmax,ymax=map(int,box)cv2.rectangle(image,(xmin,ymin),(xmax,ymax),(0,255,0),2)cv2.putText(image,f'{score:.2f}',(xmin,ymin-10),cv2.FONT_HERSHEY_SIMPLEX,0.9,(0,255,0),2)cv2.imshow('Detected People',image)cv2.waitKey(0)cv2.destroyAllWindows()# 实际使用image_tensor=preprocess_image('path_to_your_test_image.jpg')predictions=detect_people(model,image_tensor)visualize_predictions(image_tensor,predictions)注意事项
- 数据划分:确保数据集划分为训练集、验证集和测试集。
- 超参数调整:根据实验结果调整学习率、批次大小等超参数。
- 模型优化:考虑使用更高级的技术如多模态融合、迁移学习等提高模型性能。
通过上述步骤,你可以构建一个用于矿场矿车坐人检测的深度学习模型,并对其进行训练和评估。