基于 FG-CLIP 多模态模型的开放词汇检测
结合多模态模型的开放词汇目标检测实现通过图文特征深度对齐实现零样本 / 小样本类别无关检测无需重新训练主干网络即可快速适配新类别目标兼顾检测精度与工程实用性。本项目源码地址https://github.com/wrq147/FGDetection项目基于FG-CLIP多模态基础模型构建主干网络与图文特征提取能力引用自https://github.com/360CVGroup/FG-CLIP一、方案整体设计思路本方案的核心思想是“借力预训练大模型聚焦检测头优化”整体架构分为三大模块预训练大模型特征提取模块、适配性检测头模块、改进型损失函数模块同时配套设计了多尺度训练策略和高效数据集处理流程。具体设计逻辑如下1. 特征提取选用预训练的视觉-语言大模型如定制化的 fgmodel复用其经过海量数据训练的视觉特征提取能力和文本特征编码能力避免从零开始训练特征提取器大幅减少训练数据需求和训练时长。2. 检测头适配设计轻量化的 AdaptedDetectHead 检测头将大模型输出的视觉特征和文本特征进行融合分别预测目标边界框和类别置信度实现特征到检测结果的高效映射。3. 损失函数优化针对目标检测中常见的正负样本不平衡、边界框回归精度低等问题改进传统 YOLO 损失函数引入 CIoU 边界框损失和 Varifocal 类别损失提升检测精度和鲁棒性。4. 训练策略优化采用多尺度训练策略动态调整输入图像尺寸增强模型对不同尺度目标的适配能力同时引入文本特征缓存机制避免重复计算提升训练效率。二、核心模块实现细节2.1 预训练大模型特征提取模块本方案选用定制化的预训练视觉-语言大模型fgmodel作为特征提取骨干该模型同时具备视觉特征提取和文本特征编码能力能够实现图像与文本的跨模态特征对齐为目标检测的类别识别提供语义支撑。在视觉特征提取方面通过模型的 get_vision_feature 方法对输入图像进行编码输出高维视觉特征维度为768该特征包含了图像的全局语义和局部细节信息能够有效表征目标的形状、纹理等特征。在文本特征编码方面通过 tokenizer 对类别名称进行分词、编码再通过 get_text_features 方法生成类别文本特征实现类别名称的语义量化。为提升训练效率本方案设计了文本特征缓存机制。首次训练时计算所有类别名称的文本特征并保存到本地class_text_feat_cache.pth后续训练时直接加载缓存的文本特征避免重复调用大模型进行文本编码大幅减少训练耗时。核心实现代码如下# 文本特征缓存逻辑cache_pathclass_text_feat_cache.pthtext_feat_cache{}ifos.path.exists(cache_path):print(✅ 找到缓存文件直接加载文本特征...)text_feat_cachetorch.load(cache_path,map_locationcpu)else:print( 未找到缓存开始计算文本特征...)withtorch.no_grad():fortxtintrain_dataset.global_class_set:cap_intokenizer(txt,paddingmax_length,max_length196,truncationTrue,return_tensorspt).to(device)featfgmodel.get_text_features(**cap_in,walk_typelong)text_feat_cache[txt]feat.cpu()torch.save(text_feat_cache,cache_path)print(f✅ 文本特征已保存到{cache_path})2.2 适配性检测头模块AdaptedDetectHead检测头是连接预训练特征和检测结果的核心模块本方案设计的 AdaptedDetectHead 模块主要实现视觉特征与文本特征的融合、边界框预测和类别置信度预测三大功能整体结构轻量化便于部署。1. 特征融合首先通过全连接层fc对视觉特征进行维度变换和归一化再与归一化后的文本特征进行矩阵乘法运算得到图像特征与各类别文本特征的相似度矩阵cls_sim。通过 logit_scale 和 logit_bias 调整相似度尺度增强特征对齐效果。同时引入置信度掩码mask将类别相似度与目标置信度结合筛选出有效视觉特征提升检测精度。2. 边界框预测采用轻量化的卷积神经网络box_head对融合后的视觉特征进行处理输出边界框预测结果。box_head 由3个卷积层组成其中包含深度可分离卷积在保证检测精度的同时减少模型参数和计算量输出维度为4对应目标的中心点坐标cx, cy和宽高w, h。3. 类别置信度预测通过 cls_head 卷积层对融合后的视觉特征进行处理输出类别置信度预测结果输出维度为1代表目标存在的置信度。同时将类别相似度矩阵cls_btm传递给损失函数用于类别损失的计算。核心实现代码如下classAdaptedDetectHead(nn.Module):def__init__(self):super().__init__()self.logit_scalenn.Parameter(torch.ones(1)*2.6592)self.logit_biasnn.Parameter(torch.zeros(1))self.fcnn.Sequential(nn.Linear(768,768),nn.LayerNorm(768),nn.GELU(),nn.Linear(768,768),)self.box_headnn.Sequential(nn.Conv2d(768,384,kernel_size3,padding1),nn.BatchNorm2d(384),nn.GELU(),nn.Conv2d(384,384,kernel_size3,padding1,groups384),nn.BatchNorm2d(384),nn.GELU(),nn.Conv2d(384,4,kernel_size1))self.cls_headnn.Sequential(nn.Conv2d(768,192,kernel_size3,padding1),nn.BatchNorm2d(192),nn.GELU(),nn.Conv2d(192,192,kernel_size3,padding1,groups192),nn.BatchNorm2d(192),nn.GELU(),nn.Conv2d(192,1,kernel_size1))defforward(self,last_hidden,text_feat):Blast_hidden.shape[0]featsizeint(last_hidden.shape[1]**0.5)fc_text_featF.normalize(text_feat,dim-1)fc_img_featself.fc(last_hidden)fc_img_featF.normalize(fc_img_feat,dim-1)# 计算图像与文本特征相似度cls_simtorch.matmul(fc_img_feat,fc_text_feat.transpose(-1,-2))cls_simcls_sim*self.logit_scale.exp()self.logit_bias cls_btmcls_sim.view(B,featsize,featsize,-1).permute(0,3,1,2).contiguous()# 生成置信度掩码max_sim,_cls_sim.max(dim-1)confidencetorch.sigmoid(max_sim)sim_meancls_sim.mean(dim-1)sim_max_smoothedsim_mean sim_max_minsim_max_smoothed.amin(dim1,keepdimTrue)sim_max_maxsim_max_smoothed.amax(dim1,keepdimTrue)sim_mean(sim_max_smoothed-sim_max_min)/(sim_max_max-sim_max_min1e-8)real_mmconfidence*sim_mean maskreal_mm.unsqueeze(-1)cls_featfc_img_feat*mask# [B, 784, 768]cls_featcls_feat.permute(0,2,1).reshape(B,-1,featsize,featsize)box_featfc_img_feat*maskfc_img_feat box_featbox_feat.permute(0,2,1).reshape(B,-1,featsize,featsize)boxself.box_head(box_feat)# [B,4,28,28]boxbox.flatten(2)# [B,4,784]cls_mapself.cls_head(cls_feat)cls_mapcls_map.flatten(2)returnbox,cls_map,cls_btm2.3 改进型 CustomYOLOLoss 损失函数传统 YOLO 损失函数存在两大痛点一是边界框回归采用普通 IoU 损失对边界框的重叠度、中心点距离和宽高比考虑不足回归精度低二是类别损失采用交叉熵损失难以解决正负样本不平衡问题。本方案针对上述问题设计了改进型 CustomYOLOLoss 损失函数包含 CIoU 边界框损失和 Varifocal 类别损失两部分。1. CIoU 边界框损失在普通 IoU 的基础上引入中心点距离DIoU和宽高比一致性CIoU两个指标能够更全面地衡量预测框与真实框的差异提升边界框回归的精度和速度。核心实现逻辑为先将中心点坐标cx, cy和宽高w, h转换为对角坐标x1, y1, x2, y2计算交集和并集得到 IoU再计算最小外接矩形的对角线平方和中心点距离平方得到 DIoU最后引入宽高比一致性系数得到 CIoU损失值为 1 - CIoU。2. Varifocal 类别损失针对正负样本不平衡问题引入 focal_alpha 和 focal_gamma 两个参数对正负样本的损失进行加权。正样本的权重由目标置信度决定负样本的权重由预测置信度和 focal_gamma 共同决定能够有效抑制负样本的干扰提升模型对正样本的关注程度。同时结合交叉熵损失进一步提升类别识别的精度。核心实现代码如下classCustomYOLOLoss(nn.Module):def__init__(self,focal_alpha0.25,focal_gamma2.0):super().__init__()self.focal_alphafocal_alpha self.focal_gammafocal_gamma self.ce_lossnn.CrossEntropyLoss()defbbox_iou_loss(self,pred_boxes,target_boxes,eps1e-7):# 解析坐标并转换为对角坐标pred_cx,pred_cy,pred_w,pred_hpred_boxes.unbind(1)tgt_cx,tgt_cy,tgt_w,tgt_htarget_boxes.unbind(1)pred_x1pred_cx-pred_w/2pred_y1pred_cy-pred_h/2pred_x2pred_cxpred_w/2pred_y2pred_cypred_h/2tgt_x1tgt_cx-tgt_w/2tgt_y1tgt_cy-tgt_h/2tgt_x2tgt_cxtgt_w/2tgt_y2tgt_cytgt_h/2# 计算交集和并集inter_x1torch.max(pred_x1,tgt_x1)inter_y1torch.max(pred_y1,tgt_y1)inter_x2torch.min(pred_x2,tgt_x2)inter_y2torch.min(pred_y2,tgt_y2)inter_wtorch.clamp(inter_x2-inter_x1,min0)inter_htorch.clamp(inter_y2-inter_y1,min0)interinter_w*inter_h pred_areapred_w*pred_h tgt_areatgt_w*tgt_h unionpred_areatgt_area-intereps iouinter/union# 计算 DIoUcwtorch.max(pred_x2,tgt_x2)-torch.min(pred_x1,tgt_x1)chtorch.max(pred_y2,tgt_y2)-torch.min(pred_y1,tgt_y1)c2cw**2ch**2eps rho2(pred_cx-tgt_cx)**2(pred_cy-tgt_cy)**2diouiou-rho2/c2# 计算 CIoUv(4/(torch.pi**2))*torch.pow(torch.atan(tgt_w/(tgt_heps))-torch.atan(pred_w/(pred_heps)),2)alphav/(1-iouveps)cioudiou-alpha*v ciou_loss1.0-cioureturnciou_loss,ioudefvarifocal_loss(self,pred,score,target):pred_sigmoidpred.sigmoid()targettarget.type_as(pred)weighttarget*score(1-target)*((1-self.focal_alpha)*pred_sigmoid.detach()**self.focal_gamma)lossF.binary_cross_entropy_with_logits(pred,target,weightweight,reductionnone)returnloss.mean()ifloss.numel()0else0.0defforward(self,pred_box,gt_box,feat_w,feat_h,pred_cls,cls_btm,gt_box_cls_indices):# 逐批次计算损失累计总损失并求平均devicepred_box.device Nfeat_w*feat_h Bpred_box.shape[0]total_ciou0total_cls0num_valid0# 生成网格坐标解码预测框ys,xstorch.meshgrid(torch.arange(feat_h,devicedevice),torch.arange(feat_w,devicedevice),indexingij)xsxs.reshape(-1).float()ysys.reshape(-1).float()xs_normxs/feat_w ys_normys/feat_h cls_btm_flatcls_btm.flatten(2)forbinrange(B):# 解码预测框坐标pbpred_box[b].permute(1,0)pcpred_cls[b].squeeze(0)dxtorch.sigmoid(pb[:,0])dytorch.sigmoid(pb[:,1])dwpb[:,2]dhpb[:,3]cxxs_normdx cyys_normdy wtorch.clamp(torch.exp(dw)*0.2,1e-5,1.0)htorch.clamp(torch.exp(dh)*0.2,1e-5,1.0)pred_decodedtorch.stack([cx,cy,w,h],dim1)gbtorch.tensor(gt_box[b]).to(device)ifgb.numel()0:continueMgb.shape[0]gt_cidstorch.tensor(gt_box_cls_indices[b]).to(device)# 筛选正样本生成目标框和类别标签pos_masktorch.zeros(N,dtypetorch.bool,devicedevice)target_boxtorch.zeros(N,4,devicedevice)cls_targettorch.zeros(N,devicedevice)cls_scoretorch.zeros(N,devicedevice)best_ioutorch.zeros(N,devicedevice)target_cls_idxtorch.zeros(N,dtypetorch.long,devicedevice)forminrange(M):cx,cy,w,hgb[m]cidgt_cids[m]gxcx*feat_w gycy*feat_h gww*feat_w ghh*feat_h radius_wmax(2.0,gw/20.5)radius_hmax(2.0,gh/20.5)# 筛选候选区域candidate_mask(torch.abs(xs-gx)radius_w)(torch.abs(ys-gy)radius_h)ifnotcandidate_mask.any():continuecand_idxtorch.where(candidate_mask)[0]# 计算候选区域的置信度和 IoUcand_scorepc[cand_idx].sigmoid()cand_boxpred_decoded[cand_idx]_,iou_candself.bbox_iou_loss(cand_box,gb[m].unsqueeze(0).expand(len(cand_idx),4))# 筛选有效候选区域keep_competeiou_candbest_iou[cand_idx]cand_idxcand_idx[keep_compete]iou_candiou_cand[keep_compete]cand_scorecand_score[keep_compete]iflen(cand_idx)0:continue# 计算对齐分数筛选 TopK 正样本img_cls_allcls_btm_flat[b]cand_cls_logitsimg_cls_all[cid,cand_idx]cls_sim_weighttorch.sigmoid(cand_cls_logits)align_scoretorch.pow(iou_cand1e-7,1.5)*torch.sqrt(cls_sim_weight1e-7)gt_areaw*h topkint(3gt_area*500)max_kmin(12,max(4,int(N/64)))topkmin(topk,max_k,len(align_score))kmin(topk,len(align_score))topk_val,topk_idxtorch.topk(align_score,k)current_idxcand_idx[topk_idx]tmp_iouiou_cand[topk_idx]fusion_labeltorch.sqrt(tmp_iou1e-7)# 更新正样本信息best_iou[current_idx]tmp_iou pos_mask[current_idx]Truetarget_box[current_idx]gb[m]target_cls_idx[current_idx]cid cls_target[current_idx]fusion_label.detach()# 计算当前批次的损失num_pospos_mask.sum().item()ifnum_pos0:ciou,iouself.bbox_iou_loss(pred_decoded[pos_mask],target_box[pos_mask])total_ciouciou.mean()cls_score[pos_mask]torch.clamp(iou.detach(),0.0,1.0)clsself.varifocal_loss(pc,cls_score,cls_target)total_clscls# 计算多类别交叉熵损失mul_cls_predcls_btm_flat[b][:,pos_mask].transpose(0,1)gt_labeltarget_cls_idx[pos_mask]mul_cls_lossself.ce_loss(mul_cls_pred,gt_label)total_clsmul_cls_loss num_valid1# 计算批次平均损失ifnum_valid0:returnpred_box.sum()*0.0total_loss(total_ciou*4.0total_cls)/num_validreturntotal_loss2.4 多尺度训练与数据集处理模块为增强模型对不同尺度目标的适配能力本方案设计了 MultiScaleBatchSampler 多尺度采样器动态调整输入图像的尺寸。采样器预设了多个符合 YOLO 标准的输入尺寸512、640、768、1024均为32的整数倍在每个批次采样完成后随机选择一个尺寸作为当前批次的输入尺寸并同步调整数据集的输入尺寸实现多尺度训练。数据集处理模块YOLODataset适配自定义数据集结构支持对图像进行缩放、填充处理同时修正边界框坐标确保预测框与输入图像尺寸对齐。针对类别数量不固定的问题设计了固定类别填充机制将每张图像的类别数量固定为40不足部分用背景类别填充保证模型输入维度的一致性。核心功能包括1. 图像预处理通过 resize_and_pad 方法将图像缩放到目标尺寸并对不足部分进行填充填充颜色为114, 114, 114避免图像变形。2. 边界框修正通过 correct_boxes 方法根据图像缩放和填充参数修正边界框的归一化坐标确保预测框与真实框的位置对齐。3. 类别填充针对每张图像的真实类别数量不足40的部分用预设的背景类别填充生成固定长度的类别名称列表和类别索引适配模型的输入要求。三、模型训练与验证流程本方案的训练流程分为训练阶段和验证阶段采用 AdamW 优化器学习率设置为1e-4权重衰减设置为1e-5训练轮次为30轮。具体流程如下1. 数据加载通过 DataLoader 加载训练集和验证集训练集采用多尺度采样器验证集采用固定批次采样。2. 模型初始化加载预训练大模型fgmodel并设置为评估模式避免训练过程中更新大模型参数初始化检测头dethead若存在预训练权重dethead_yolo_best.pth则加载部分权重实现迁移学习。3. 训练阶段逐批次加载图像和标签通过预训练大模型提取视觉特征和文本特征输入到检测头得到预测结果计算改进型损失函数通过反向传播更新检测头参数实时打印训练损失通过 tqdm 显示训练进度。4. 验证阶段每轮训练完成后切换模型为评估模式逐批次加载验证集数据计算验证损失若当前验证损失为历史最优则保存检测头权重dethead_yolo_best.pth。5. 训练完成训练结束后保存最终的检测头权重dethead_yolo_final.pth用于后续推理部署。预览图片识别文本”左边的猫“识别文本”书“识别文本”狗“

相关新闻