用SAM做红外小目标检测:冻结编码器与检测头实现
简介离线配套资源面向红外小目标检测方向的研究者与工程技术人员解决低对比度、复杂背景下弱小目标难以分割与识别的问题算法核心为统计区域合并SAM。包内含53个文件24个Python脚本涵盖数据集读取、预处理、SAM模型构建、训练与评估等环节22个pyc为编译缓存6个txt为运行说明另有1个Markdown文档介绍整体结构压缩包仅145KB轻量便于快速部署。已有95人学习下载。源码可直接运行并内置Sirstv2、IRSTD-1k、NUDT-SIRST等红外数据集的处理逻辑方便对比不同尺度、不同场景下的检测效果。项目源码结构清晰通过阅读与调试能够掌握SAM在红外图像分割中的应用流程理解多尺度处理、深度网络融合等改进思路为军事侦察、航空航天、智能交通等实际项目提供可复用的参考实现。1. 用SAM做红外小目标检测先放下“大模型不配做小目标”的偏见红外小目标检测有一个经常被低估的事实目标在画面里往往只有几个到十几个像素拿YOLO这类通用检测器去跑漏检率能把人气到换方向而换用SAM这套为通用分割设计的大模型反而能在低信噪比的场景里拿出更稳的特征表达。这份项目源码做的正是基于SAM的红外小目标检测算法实现——不拿SAM直接出框而是把它当成特征底座接一个专门针对小目标设计的检测头在NUAA-SIRST这类公开红外小目标检测数据集上能完整复现训练和推理。适合想快速拿到一条可改、可续训、可落地的红外小目标检测基线的人也适合正在做SAM相关应用的开发者用来对比特征效果。2. 为什么是SAM而不是换个YOLO红外小目标的难点与接入路径选型先别急着打开代码得先把“红外小目标难在哪”和“SAM凭什么能解”这两件事说透。方向选错了后面改代码全是白费劲。2.1 红外小目标到底“小”在哪9×9像素、低信噪比、无纹理红外小目标检测IRSTD常用于无人机反制、制导告警、海上搜救这类前视红外场景。这里的“小”不是相对小而是绝对小公开数据集里成熟的定义是目标尺寸小于图像面积的0.12%~0.15%换算到512×512的图上通常就是9×9像素以内极端情况下只有2~4个像素和一个噪点差不多。加上红外图像是单通道灰度目标没有颜色、几乎没有纹理和云层边缘、海面杂波、地物高亮在视觉上高度相似信噪比经常低到目标本身都看不清。传统做法围绕“局部对比度”做文章。Top-hat形态学滤波、LCM局部对比度方法、以及基于IPI的稀疏低秩分解核心都在赌一个前提目标比邻域亮且比邻域小。这个前提在纯净天空下成立在复杂地物背景下就扛不住云层边缘和建筑尖角带来的高亮伪目标会把虚警率拉得很高。深度学习进来之后用CNN做分割式检测成了主流但通用检测器在红外小目标上的表现并不理想原因在于方法类型代表优点短板形态学Top-hat实现简单、实时复杂背景虚警多局部对比LCM对孤立点目标有效参数依赖场景稀疏低秩IPI理论清晰速度慢、大目标失效CNN检测YOLO系列端到端、速度快小目标样本在anchor里占比太低分割类UNet/FCN像素级输出需要像素级标注2.2 SAM的原理和它真正的价值不是分割而是编码器SAMSegment Anything Model是Meta在2023年开源的提示分割模型结构由三块组成图像编码器一个ViT把输入图像编码成高维特征、提示编码器把点、框、掩膜转成embedding、掩膜解码器把二者组合后输出分割掩膜。它本身解决的不是“检测”而是“你给我一个提示我还你一个分割结果”。很多人把SAM算法当成一个黑盒分割器来用直接拿官方权重去分割红外小目标结果一塌糊涂然后得出“SAM不适用于小目标”的结论。这个判断只对了一半。SAM真正值钱的部分是它的图像编码器ViT的全局注意力机制天然能看到整幅图的上下文红外场景里“目标极亮、邻域极暗”这种局部突变在原图尺度可能只有几个像素但在Transformer的全局关系里是一个显著的离群点。更重要的是SAM是在SA-1B这类超大规模数据上预训练过的它的特征表达对低纹理、低语义的目标比从零训练的CNN稳定得多。很多IRSTD方向的工作直接用SAM的编码器特征灌进一个小检测头就能超过专门设计的小目标网络这在第一次看到时是有点反直觉的。2.3 接入SAM的三种路径直接推理、冻结特征、微调要回答“怎么用SAM”得先把三条路摆出来对比。第一条路是直接推理给一个框提示或点提示让SAM把目标分割出来再取外接框。这个思路在红外小目标上会当场翻车——SAM的图像编码器在1024×1024输入下做了16倍下采样特征图是64×64一个9像素的目标在特征图上连1个像素都占不满分割掩膜基本是噪声。且官方权重不擅长处理16bit红外灰度图直接推理的分割质量极差。第二条路是项目源码采用的把SAM的图像编码器当作冻结的backbone重新设计检测头。这条路的好处有三个。其一SAM已经在大规模自然图像上收敛得很充分不需要大量红外数据去从头训其二冻结状态下显存和训练成本都降一个量级单卡能跑其三检测头可以完全围绕小目标特性设计不受SAM原本分割逻辑约束。第三条路是SAM-Adapter或LoRA这类微调方案效果上限更高但需要额外的训练预算和调参适合有标注数据、有算力的后续迭代。接入方式训练成本小目标表现适用场景SAM直接推理零差下采样丢目标快速试跑、大目标冻结编码器检测头低好推荐首选LoRA/SAM-Adapter微调中更好有标注、有算力我一般会建议先走冻结编码器检测头这条路因为它能最快验证“SAM对这个数据集到底有没有用”。项目源码默认给的也是这条配置后面的章节顺着这个主线讲。3. 项目结构与核心模块拿到源码先看这四个文件下载下来的项目源码是一个完整的工程包不是零散脚本拼凑的demo。把它当工程看第一步不是急着跑train.py而是先搞清楚骨架。3.1 目录结构先分清哪里是骨架哪里能改源码包打开通常是这样的布局IRSTD-SAM/ ├── configs/ │ └── train.yaml # 训练参数总入口 ├── dataset/ │ ├── base.py # 数据读取与预处理 │ └── nuaa_sirst.py # NUAA-SIRST专用数据类 ├── models/ │ ├── sam_backbone.py # SAM编码器加载与封装 │ ├── head.py # 检测头 │ └── irstd_sam.py # 主网络拼装 ├── tools/ │ ├── train.py # 训练入口 │ ├── test.py # 测试/推理入口 │ └── visualize.py # 检测结果可视化 ├── weights/ │ └── sam_vit_b.pth # SAM官方预训练权重 └── README.md # 运行说明其中models/irstd_sam.py是主网络的拼装文件models/sam_backbone.py是SAM的适配层configs/train.yaml是全部超参数。我个人会把四个文件先读一遍irstd_sam.py、sam_backbone.py、head.py、train.yaml。其它文件基本都是数据IO和训练循环出了问题再回头看。3.2 SAM编码器如何封装成backbone冻结权重和eval模式是关键第2章说的“冻结编码器检测头”落到代码里其实就是把官方SAM拆开只保留图像编码器。常见做法是直接用官方仓库的sam_model_registry加载完整权重然后取.image_encoder子模块# models/sam_backbone.py from segment_anything import sam_model_registry import torch.nn as nn class SAMBackbone(nn.Module): def __init__(self, ckpt_path, freezeTrue): super().__init__() # 完整加载SAM权重再取出图像编码器 sam sam_model_registry[vit_b](checkpointckpt_path) self.encoder sam.image_encoder self.out_channels 256 if freeze: for p in self.encoder.parameters(): p.requires_grad False self.encoder.eval() # 冻结时切换为eval稳定LayerNorm行为 def forward(self, x): # SAM编码器输出形状: (B, 256, H/16, W/16) return self.encoder(x)这段代码里有几个容易忽略的点。sam_model_registry[vit_b]要求传入官方checkpoint路径它会自动构建完整的SAM模型并把权重填进去取.image_encoder之后prompt_encoder和mask_decoder就不再参与前向。out_channels设为256因为ViT-B的图像编码器最终输出通道数是256。冻结时调用eval()是为了让LayerNorm和Dropout的行为在训练和推理保持一致虽然ViT里没有BN但eval这个习惯对任何冻结模块都是安全的。forward返回的特征图分辨率是输入的1/16比如输入512×512特征图是32×32。这就是第2章提到的“小目标在特征图上占不满一个像素”的问题所在解决办法在4.3节的数据增强里展开说在第5章避坑。3.3 检测头CenterNet风格的无锚框回归为什么检测头不选YOLO的anchor方案红外小目标本质上是点目标一个9×9的目标在32×32特征图上只是一个亮点。anchor方案要预先铺几千个先验框正样本比例极低训练时损失几乎被背景淹没。CenterNet风格的无锚框方案更契合预测一张中心点热力图热力图的峰值就是目标中心再回归目标宽高和中心点偏移。这个思路在红外小目标检测里几乎是标配代码也干净# models/head.py import torch.nn as nn import torch.nn.functional as F class IRSTDHead(nn.Module): def __init__(self, in_ch256, num_cls1): super().__init__() # 用3x3卷积把SAM的256维特征压缩到128维减少后续计算量 self.reduce nn.Sequential( nn.Conv2d(in_ch, 128, 3, padding1, biasFalse), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue) ) self.cls_head nn.Conv2d(128, num_cls, 1) # 中心点热力图 self.reg_head nn.Conv2d(128, 4, 1) # 目标框 (x,y,w,h) 归一化 def forward(self, x): feat self.reduce(x) heatmap self.cls_head(feat) box self.reg_head(feat).sigmoid() return heatmap, boxreduce层的作用是把256通道压缩到128主要省显存和时间。cls_head输出形状是(B,1,H,W)每个位置是“该位置是目标中心的置信度”reg_head输出4个值分别是中心相对该网格的偏移和目标的宽与高最后sigmoid把输出压到0~1之间推理时乘回输入尺寸。训练时热力图标签不是二值点而是以目标中心为中心的高斯核半径通常取3像素这样中心点附近的像素也能参与损失计算梯度更平滑。3.4 损失函数Focal Loss怎么压住背景主导的训练红外小目标训练时最麻烦的问题一张512×512的图里往往只有1~2个目标正样本可能只有几十个像素背景有几十万像素。如果用普通的BCE损失模型很快会学会“全部预测为背景”因为那样损失已经很低。Focal Loss在这里是刚需# models/loss.py import torch def ir_focal_loss(pred, target, alpha2.0, beta4.0): pred pred.clamp(min1e-6, max1 - 1e-6) pos (target 1.0).float() neg (target 1.0).float() # 正样本让预测接近1负样本高斯核外围离中心越远惩罚越小 pos_loss (1 - pred) ** alpha * torch.log(pred) * pos neg_loss (1 - target) ** beta * pred ** alpha * torch.log(1 - pred) * neg n_pos pos.sum().clamp(min1) return -(pos_loss.sum() neg_loss.sum()) / n_posalpha和beta是Focal Loss的两个旋钮。alpha控制难易样本的权重设为2.0时预测置信度高的负样本被压得很低模型会把梯度集中在难分类的样本上beta是专门为CenterNet风格热力图准备的target是高斯核生成的连续值离中心越远的像素越接近0这些点虽然算负样本但因为(1-target)^beta这个系数内核边缘的像素惩罚更轻避免把中心点周围学得太“干净”而导致中心定位偏移。总损失一般再加一个简单的L1框损失权重配比在train.yaml里给的是10:1热力图损失占大头。4. 从数据集到推理训练与复现的完整操作流结构看完就该动手复现了。这一章按“数据准备→预处理→训练参数→推理”的顺序走每一步都给能直接抄的命令。4.1 数据集准备NUAA-SIRST和IRSTD-1k的目录约定红外小目标检测数据集目前比较常用的是NUAA-SIRST和IRSTD-1k。前者规模小单帧目标少适合快速验证后者目标数量更多、背景更杂适合评估模型上限。它们的目录结构基本是“图像掩膜”对排布NUAA-SIRST/ ├── images/ │ ├── 0001.png │ ├── 0002.png │ └── ... ├── masks/ │ ├── 0001.png │ ├── 0002.png │ └── ... └── trainval.txt # 每一行一个样本名不带扩展名masks里的标注是二值图目标像素为255背景为0。拿到这份项目源码后我习惯先把trainval.txt读进来确认图片和掩膜能一一对上再跑数据集类做冒烟测试。很多复现失败其实是路径写错或者掩膜名称和图像名不对应导致的这种问题定位起来比训练bug还烦人。4.2 红外单通道的三通道化与归一化吃透预处理再说训练红外相机输出的通常是单通道灰度图常见做法是读成灰度图后复制三份再去适配SAM预训练权重。这里有个细节SAM官方权重是在自然图像上预训练的自然图像是8bit RGB而红外图有的相机输出是14bit甚至16bit。如果不做灰度拉伸直接除以65535图像整体偏暗编码器特征会偏离预训练分布。# dataset/base.py import cv2 import numpy as np def load_ir_image(path, size(512, 512)): # 以灰度模式读取红外图兼容16bit数据 img cv2.imread(path, cv2.IMREAD_UNCHANGED) if img is None: raise FileNotFoundError(fcannot read {path}) if img.dtype ! np.uint8: # 16bit/12bit数据先做min-max拉伸到8bit再走后续流程 img cv2.normalize(img, None, 0, 255, cv2.NORM_MINMAX) img cv2.resize(img, size, interpolationcv2.INTER_LINEAR) img img.astype(np.float32) / 255.0 # 单通道复制为三通道匹配SAM的RGB输入约定 img np.stack([img, img, img], axis-1) return img三个关键点。一是用IMREAD_UNCHANGED读图避免cv2把16bit图截断成8bit二是用NORM_MINMAX做整图拉伸而不是固定除以65535因为不同场景的整体亮度差异极大最大值经常被离群亮点带偏三是复制成三通道SAM的ViT预训练结构里第一层卷积是3通道输入直接喂灰度图会报维度错误。4.3 训练参数表与调参策略项目源码的configs/train.yaml里给的默认参数是一套经过验证的配置单卡就能跑。我实际复现时会在它基础上按显存微调参数默认值说明input_size512×512太小目标消失太大显存吃不消batch_size83090/4090可用16G显存建议降到4epochs100冻结backbone时60~80轮基本收敛optimizerAdamWlr1e-4weight_decay1e-4schedulercosine配合warmup 5轮前几轮loss更稳backbone_freezetrue冻结SAM编码器heatmap_loss_weight10.0热力图损失权重box_loss_weight1.0框回归损失权重训练命令走tools/train.pypython tools/train.py \ --config configs/train.yaml \ --data-root /data/NUAA-SIRST \ --batch-size 8 \ --lr 1e-4 \ --gpu 0--data-root必须指向包含images和masks的父目录--batch-size是总batch多卡时会被自动平分。lr建议从1e-4起步不要调大因为检测头只有三层卷积学习率过大直接震荡。如果显存不够优先降batch_size到4然后加梯度累积。训练日志会每个epoch输出一次mIoU、precision和recall我习惯盯着recall看因为红外小目标场景漏检的代价比误检高得多。4.4 推理与可视化峰值提取和demo输出推理阶段的核心是把热力图的局部最大值解析成目标框。test.py内部先用前向拿到heatmap和box然后做峰值提取# tools/test.py 的核心片段 import numpy as np def peak_extract(heat, topk100, thresh0.5): # heat: (H, W) 目标中心热力图 flat heat.ravel() idx np.argpartition(flat, -topk)[-topk:] xs, ys np.unravel_index(idx, heat.shape) scores flat[idx] keep scores thresh xs, ys, scores xs[keep], ys[keep], scores[keep] # 按置信度降序排列 order np.argsort(-scores) return list(zip(xs[order], ys[order], scores[order]))argpartition是取前topk个最大值的常用技巧比全量排序快很多thresh控制置信度下限默认0.5。提取到中心点后按reg_head输出的宽高换算成框坐标画在图上保存。调试时我会把输入图、热力图、检测结果三张图并排输出确认中心点落位是否和掩膜中心一致。可视化命令是python tools/visualize.py --ckpt checkpoints/irstd_sam.pth \ --dir /data/NUAA-SIRST/images --out ./vis注意train.py和test.py都把配置文件当作唯一超参数入口命令行参数只做覆盖所以如果改了train.yaml里的输入尺寸推理时要保持一致否则中心点坐标按错缩放比还原。提示如果用的是IRSTD-1k这类大规模数据集建议先把trainval.txt拆成train.txt和val.txt项目默认按8:2划分别让验证集混进训练。5. 复现避坑五个最容易翻车的点与排查方案复现文稿这种事跑通一次不难难的是中间那些“看起来没报错但结果就是不对”的环节。下面五条全部来自实际踩坑按“现象→原因→解决”写。5.1 训练loss不降反升前几轮直接爆掉现象第一个epoch结束loss比初始还高热力图输出全是0或者全是背景。原因绝大多数是输入分布没对齐——红外图直接除以65535或者忘记做灰度拉伸喂给SAM的值域和ImageNet预训练时差太多ViT的LayerNorm统计量被输入分布带偏。解决严格按照4.2节的NORM_MINMAX先拉伸到0~255再除以255之后用torchvision的Normalize补充ImageNet的mean和std。另一个坑是热力图最后一层卷积初始化全为0导致初始输出全是噪声给cls_head的bias初始化为-4.0即可稳定训练。5.2 训练中显存溢出batch直接OOM现象batch_size设8输入512×512跑到第二个epoch报CUDA out of memory。原因很多人只算了检测头的显存忘了SAM编码器虽然冻结但前向依然占用显存ViT-B在512×512输入下骨干就要吃掉接近5~6GB。解决16G显存的卡把batch_size降到4或者把输入降到448×448更稳的做法是打开梯度累积batch_size4加累积2步等效batch为8训练曲线不会差太多。输入尺寸不要低于384否则9像素的小目标在16倍下采样后彻底变成亚像素再好的检测头也无济于事。5.3 加载官方sam_vit_b.pth时报错KeyError现象把sam_vit_b.pth直接load到自定义的SAMBackbone.encoder里报unexpected key或者size mismatch。原因官方checkpoint是以完整SAM模型为粒度保存的state_dict里的键都带image_encoder、prompt_encoder、mask_decoder前缀而如果只建了一个encoder子模块去load键名对不上是必然的。解决不要自己组装用官方sam_model_registry vit_b 加载完整权重再从sam.image_encoder取出编码器像3.2节那样。如果编码器外面又包了一层自定义类比如加了一个降维卷积load_state_dict时加strictFalse只挑匹配键。5.4 3×3小目标检测不到漏检全集中在最小目标现象5×5以上的目标都能检出3×3以下几乎全军覆没热力图上对应位置没有任何响应。原因这是下采样和标签高斯核半径共同造成的。512输入经过16倍下采样后特征图只有32×32一个3像素目标对应特征图0.19像素而训练时高斯核半径设3等于把中心附近算成多个正样本3像素目标参与损失的有效像素太少模型根本没学到它的响应模式。解决第一输入提到640×640让目标在特征图上至少占1个像素第二高斯核半径随目标尺寸自适应目标直径小于5像素时半径改为1~2第三训练时对包含小目标的原图做2倍随机放大裁剪等价于给小目标上采样。5.5 太阳、尾焰这类高亮区域被当成目标现象误检集中在图像里最亮的几块区域置信度甚至比真目标还高。原因SAM编码器擅长“抓显著物”红外场景里亮度最高的区域往往不是目标而是背景干扰源单纯把关口设在热力图置信度上挡不住这类“外观即显著”的干扰。解决在训练数据里显式加入难负例。做法是从不包含目标的红外片段里随机裁出高亮斑块直接拼到训练图的背景区域并保证这些位置的热力图标签为0另外可以在预处理阶段给模型额外输入一个局部对比度通道——取目标级别的小窗口7×7计算中心像素与邻域均值的差分作为第4个输入通道这样模型会学到“目标不仅亮而且比周围亮一截”的特征。6. 进阶用法用SAM的分割能力给检测结果兜底验证6.1 检测框转分割掩膜的三行代码检测头只输出中心点和宽高验证时拿到的只是框。有些场景需要像素级目标区域比如统计目标面积、形态或者给标注工具半自动生成高精度掩膜。这时可以把检测框当作提示喂给完整版SAM让它在原图上做一次精分割from segment_anything import sam_model_registry, SamPredictor sam sam_model_registry[vit_b](checkpointweights/sam_vit_b.pth) predictor SamPredictor(sam) predictor.set_image(rgb_input) # 红外三通道图尺寸对齐 mask, score, _ predictor.predict( boxbox_xyxy, # 检测框坐标xyxy格式 multimask_outputFalse # 只取最优掩膜 )这个流程的适用范围需要注意检测框足够准的时候SAM补出来的掩膜质量很高能直接当伪标注用但目标如果小于5×5像素SAM的mask会碎成斑点。我一般会把检测框外扩1~2个像素再喂进去让提示信息更充分得到的掩膜反而更完整。6.2 把伪标注滚雪球回训练集进阶玩法是把这套流程变成一个半自动标数据管线先用现有权重在无标注红外切片上推理拿到检测框和置信度再框选高置信结果跑一次上面代码得到掩膜人工只需快速过滤明显错检标注效率能提一个量级。对新增场景往往手动标几百帧就够让模型继续迭代一轮。我从一开始做这个项目时也怀疑过——SAM一个为自然图像分割设计的模型凭什么管红外小目标的事但代码跑通之后我养成了一个习惯拿到任何新红外数据集先做一次“检测框→SAM精分割”的闭环验证既看检测有没有漏也顺手攒一批高置信伪标注。这个习惯帮我避开了不少盲调参数的弯路。这份源码包里的配置和脚本就是按这条路径整理的需要的直接把项目源码下载下来按第4章的目录组织数据就能跑通。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联
返回资讯列表 →