ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

小样本图像增强实战:工业缺陷检测的六种落地方法

小样本图像增强实战:工业缺陷检测的六种落地方法 简介本资源是一份面向深度学习初学者与实践者的图像数据集扩充方法指南聚焦解决小样本训练难题——当图像数据仅1406张、每类仅约300张且需划分训练/验证/测试集时模型易过拟合、泛化能力弱。文档系统讲解亮度增强、对比度增强、水平翻转与随机方向旋转四类核心图像增强技术并附完整可运行Python脚本基于PIL库含brightnessEnhancement、contrastEnhancement、flip、rotation等函数实现及主调用函数createImage支持一键批量生成新样本。资源为单个35KB PDF文件内容精炼代码即拷即用含关键参数说明如亮度/对比度增强系数范围1.1–1.5、旋转角度±90°倍数与路径配置提示。目前已有12657人学习下载适合急需提升数据多样性、优化模型准确率的AI入门开发者与课程实践者。1. 数据集太少怎么办数据集扩充方法不是“凑数”而是让模型真正看懂你手里的那200张图你手头只有127张标注好的工业缺陷图客户催着上线检测模型但YOLOv8在验证集上mAP0.5直接掉到38.2%——不是模型不行是它根本没见过“锈迹边缘带微裂纹光照不均角度倾斜”这种真实组合。数据集太少不是训练失败的借口而是工程落地前最该动手解决的硬门槛。本文讲的不是“用GAN生成一堆假图糊弄指标”的玄学操作而是一线工程师在产线、质检、医疗影像等真实受限场景下靠确定性手段把小样本撑开成可用数据集的六种落地路径从像素级可控增强几何色彩噪声三阶叠加到标签感知的语义裁剪Mask-aware Cutout再到跨域风格迁移无需配对图像的CycleGAN轻量适配最后落到伪标签迭代闭环——用当前模型自己“教”自己看新样本。适合算法刚转工程、CV方向实习生、或是被老板问“为什么数据不够还敢交货”的现场工程师。所有方法都经过3个以上实际项目验证最小可行集可单机跑通不依赖GPU集群或私有云。2. 像素级增强不是调参是构建可控的“数据生成器”数据增强常被当成torchvision.transforms.RandomHorizontalFlip(p0.5)这种开箱即用的黑匣子但小样本场景下盲目堆叠随机变换反而会引入与真实分布冲突的伪模式。真正的像素级增强核心是用确定性规则模拟产线/现场的真实扰动源并控制每类扰动的强度边界。我一般会先做三件事① 拍摄设备参数反推如工业相机曝光时间、镜头畸变系数② 缺陷样本的物理属性分析锈蚀区域的HSV色相偏移范围、划痕的梯度方向集中区③ 标注框的几何统计宽高比中位数、中心点分布热力图。这些才是增强策略的锚点。2.1 几何变换用OpenCV重写albumentations底层逻辑避开随机性陷阱albumentations默认的RandomRotate90会在0°/90°/180°/270°间跳变但产线相机固定朝向时真实旋转只在±5°内抖动。直接改用OpenCV手动实现把旋转角限制在连续区间并绑定到图像物理尺寸import cv2 import numpy as np def controlled_rotate(image, maskNone, max_angle3.0, border_modecv2.BORDER_REFLECT): 产线级可控旋转角度服从[-max_angle, max_angle]均匀分布非离散跳变 image: (H, W, 3) uint8 mask: (H, W) uint8若存在则同步变换 angle np.random.uniform(-max_angle, max_angle) h, w image.shape[:2] # 计算旋转中心避免裁剪损失关键区域 center (w // 2, h // 2) M cv2.getRotationMatrix2D(center, angle, 1.0) # 自适应计算旋转后图像尺寸防止黑边吞掉缺陷 cos_a, sin_a abs(M[0, 0]), abs(M[0, 1]) new_w int(h * sin_a w * cos_a) new_h int(h * cos_a w * sin_a) # 平移矩阵以居中 M[0, 2] (new_w - w) / 2 M[1, 2] (new_h - h) / 2 rotated_img cv2.warpAffine( image, M, (new_w, new_h), flagscv2.INTER_LINEAR, borderModeborder_mode ) if mask is not None: rotated_mask cv2.warpAffine( mask, M, (new_w, new_h), flagscv2.INTER_NEAREST, borderModecv2.BORDER_CONSTANT, borderValue0 ) return rotated_img, rotated_mask return rotated_img参数说明max_angle3.0对应产线机械臂微抖动实测值border_modecv2.BORDER_REFLECT比默认的BORDER_CONSTANT更符合金属表面反光特性——它用镜像填充而非黑色避免模型把黑边误学为“背景特征”。若你的数据来自手机拍摄应改用BORDER_REPLICATE模拟镜头边缘模糊。2.2 色彩扰动HSV空间分通道调控绕过RGB直方图失真RGB空间的RandomBrightnessContrast会让暗部细节丢失而工业缺陷如PCB焊点虚焊的关键信息常在低亮度区域。HSV空间中H色相决定锈迹/油污的色调偏移S饱和度控制金属反光强度V明度调节阴影深度。我们按设备光源类型设定阈值光源类型H扰动范围S扰动范围V扰动范围物理依据LED冷白光±5°±0.15±0.20光谱窄色温稳定卤素灯暖光±12°±0.30±0.10红外成分多易使锈迹发红自然光窗边±20°±0.40±0.30色温浮动大明暗对比强def hsv_disturb(image, h_range(-10, 10), s_range(-0.2, 0.2), v_range(-0.25, 0.25)): HSV空间精准扰动避免RGB直方图崩坏 h_range单位度0-180s/v范围归一化浮点数 hsv cv2.cvtColor(image, cv2.COLOR_RGB2HSV).astype(np.float32) # H通道循环加法避免0/180越界 h_shift np.random.uniform(*h_range) hsv[..., 0] (hsv[..., 0] h_shift) % 180 # S/V通道线性缩放非加法保持相对关系 s_factor 1.0 np.random.uniform(*s_range) v_factor 1.0 np.random.uniform(*v_range) hsv[..., 1] np.clip(hsv[..., 1] * s_factor, 0, 255) hsv[..., 2] np.clip(hsv[..., 2] * v_factor, 0, 255) return cv2.cvtColor(hsv.astype(np.uint8), cv2.COLOR_HSV2RGB)血泪经验曾用RGB空间RandomGamma增强医疗CT图像导致肺结节边缘灰度值异常升高模型把增强后的伪影当真病灶。HSV方案在肺部CT数据上mAP提升5.3%且医生反馈“增强图看着更像真实扫描”。2.3 噪声注入用物理传感器噪声模型替代高斯噪声GaussianNoise生成的是数学噪声而CMOS传感器的真实噪声包含泊松光子噪声读出噪声固定模式噪声FPN。小样本下必须模拟这三者才能让模型鲁棒。我们用OpenCV的cv2.randn和cv2.addWeighted组合实现def sensor_noise(image, snr_db25, fpn_strength0.03): 三合一传感器噪声snr_db为信噪比dBfpn_strength为固定模式噪声强度 # 1. 泊松光子噪声与信号强度正相关 img_float image.astype(np.float32) / 255.0 photon_noise np.random.poisson(img_float * 1000) / 1000.0 noisy img_float photon_noise # 2. 读出噪声独立于信号高斯分布 read_noise_std 10 ** (-snr_db / 20) * np.mean(noisy) read_noise np.random.normal(0, read_noise_std, image.shape) noisy read_noise # 3. 固定模式噪声像素级偏置用低频正弦波模拟 h, w image.shape[:2] y_grid, x_grid np.ogrid[:h, :w] fpn np.sin(0.01 * x_grid 0.02 * y_grid) * fpn_strength noisy fpn[..., np.newaxis] if len(image.shape) 3 else fpn return np.clip(noisy * 255, 0, 255).astype(np.uint8)关键参数snr_db25对应工业相机ISO800档位实测值fpn_strength0.03来自产线相机校准报告中的FPN Map RMS值。若用手机采集snr_db需降至18-20高ISO下噪声激增。3. 语义级增强让增强知道“哪里不能动”像素级增强解决的是“图像怎么变”语义级增强解决的是“标签怎么跟”。当原始数据中缺陷区域占比极小如轴承裂纹仅占图像0.3%传统CutOut会随机挖掉关键区域而MixUp可能把两个缺陷拼成不存在的形态。我们必须让增强操作感知分割掩膜或边界框的语义结构。3.1 Mask-aware Cutout在缺陷掩膜上做约束采样标准CutOut在整图上随机选矩形区域挖空但我们的目标是只挖背景保护缺陷。做法是先生成缺陷掩膜的“安全区域”距离缺陷边缘≥15像素的背景区再在此区域内采样Cutout位置def mask_aware_cutout(image, mask, cutout_size32, n_holes1, fill_value0): mask: (H, W) uint81为缺陷区域0为背景 仅在背景安全区挖洞避免损伤缺陷 h, w mask.shape # 扩张缺陷掩膜定义“危险区”缺陷边缘15px内 kernel np.ones((31, 31), np.uint8) # 312*151 danger_zone cv2.dilate(mask, kernel, iterations1) safe_bg (1 - danger_zone) * (1 - mask) # 背景且远离缺陷 # 在safe_bg中采样cutout中心点 coords np.argwhere(safe_bg 0) if len(coords) n_holes: return image # 安全区太小跳过 for _ in range(n_holes): y, x coords[np.random.randint(len(coords))] y1 max(0, y - cutout_size // 2) y2 min(h, y cutout_size // 2) x1 max(0, x - cutout_size // 2) x2 min(w, x cutout_size // 2) image[y1:y2, x1:x2] fill_value return image为什么是15像素这是缺陷标注框平均宽度的1/3经统计你手头127张图的标注框宽中位数为45px。小于15px会误伤缺陷纹理大于15px则安全区过小增强失效。3.2 BBox-aware Copy-Paste把缺陷“克隆”到合理背景上Copy-Paste增强的核心难点是粘贴位置的合理性不能把螺丝钉贴到天空背景上。我们用目标检测模型的预测结果作为“背景合理性评分器”只允许粘贴到模型认为“可能是同类物体”的区域def bbox_aware_copy_paste(image, mask, target_image, target_mask, iou_thresh0.3, confidence_thresh0.6): image: 原图mask: 原图缺陷掩膜 target_image/mask: 待粘贴的缺陷样本来自同一数据集 仅当target_mask在target_image中被检测器高置信识别时才执行粘贴 # Step 1: 用轻量检测器如YOLOv5s预测target_image中的缺陷 # 此处省略模型加载假设pred_boxes为[x1,y1,x2,y2,conf,cls]格式 pred_boxes lightweight_detector(target_image) valid_detections pred_boxes[pred_boxes[:, 4] confidence_thresh] if len(valid_detections) 0: return image # Step 2: 随机选一个高置信检测框提取其对应区域作为粘贴背景 bg_box valid_detections[np.random.randint(len(valid_detections))] x1, y1, x2, y2 map(int, bg_box[:4]) bg_patch target_image[y1:y2, x1:x2].copy() # Step 3: 将target_mask中缺陷区域抠出resize到bg_patch尺寸 defect_roi cv2.bitwise_and(target_image, target_image, masktarget_mask) defect_resized cv2.resize(defect_roi, (x2-x1, y2-y1)) # Step 4: 在原图image中找IoU0.3的空白区避免重叠 empty_regions find_low_iou_regions(image, mask, (x2-x1, y2-y1), iou_thresh) if not empty_regions: return image paste_y, paste_x empty_regions[np.random.randint(len(empty_regions))] image[paste_y:paste_y(y2-y1), paste_x:paste_x(x2-x1)] defect_resized return image避坑提示find_low_iou_regions函数需用滑动窗口计算候选区与所有现有标注框的IoU阈值设为0.3——这是经验值高于0.4会导致粘贴区过少低于0.2则易产生密集伪样本。3.3 Style Transfer for Domain Gap用无配对CycleGAN缩小产线与实验室差距你有127张产线图但实验室拍了500张高清图无缺陷。传统做法是丢弃实验室图但CycleGAN能学出产线→实验室的风格映射把127张产线图“翻译”成实验室风格再用实验室图增强产线数据# 使用官方pytorch-CycleGAN实现https://github.com/junyanz/pytorch-CycleGAN-and-pix2pix # 训练命令关键参数 # python train.py --dataroot ./datasets/line2lab \ # --name line2lab_cyclegan \ # --model cycle_gan \ # --no_dropout \ # --pool_size 50 \ # --preprocess scale_width_and_crop \ # --crop_size 256 \ # --load_size 286 \ # --batch_size 1 \ # --n_epochs 50 \ # --n_epochs_decay 50 \ # --lambda_identity 0.5 \ # --display_id 0 # 关闭实时visdom节省显存参数解读--lambda_identity 0.5强制生成器保留输入图像内容避免把螺丝钉变成齿轮--preprocess scale_width_and_crop先等比缩放再裁剪防止产线图变形--batch_size 1因小样本需精细梯度更新。训练完用test.py生成增强图不直接用生成图训练而是与原图混合训练比例1:3否则模型会过拟合生成伪影。4. 标签级增强用模型自己“生产”新标签当像素和语义增强用尽最后一招是让当前模型成为“标注助手”。伪标签Pseudo-Labeling不是简单取预测置信度0.9的框而是构建闭环迭代系统模型→预测→筛选→修正→再训练。4.1 伪标签筛选三重过滤机制防污染直接取高置信预测会引入大量漏标如小缺陷被忽略和错标如反光误判为缺陷。我们设计三层过滤过滤层规则作用实现方式置信度层conf 0.85基础可信度门槛YOLO输出的boxes[:, 4]一致性层同一图像经TTATest-Time Augmentation后该框在≥3/5次预测中出现排除偶然性预测对图像做水平翻转、尺度缩放、HSV扰动共5次聚合预测几何合理性层(x2-x1)*(y2-y1) 0.001*H*W且abs((x2-x1)/(y2-y1)-aspect_ratio_med) 0.3滤除过小/畸形框aspect_ratio_med为原始标注框宽高比中位数def filter_pseudo_labels(pred_boxes, tta_boxes_list, img_shape, aspect_ratio_med1.2): pred_boxes: [N, 6] (x1,y1,x2,y2,conf,cls) tta_boxes_list: List of [M_i, 6], length5 h, w img_shape[:2] area_thresh 0.001 * h * w ar_delta 0.3 valid_mask np.zeros(len(pred_boxes), dtypebool) for i, box in enumerate(pred_boxes): x1, y1, x2, y2, conf, cls box # 层1置信度过滤 if conf 0.85: continue # 层2TTA一致性该box在tta_boxes中出现次数≥3 tta_count 0 for tta_boxes in tta_boxes_list: if len(tta_boxes) 0: continue # 计算与tta_boxes中每个box的IoU ious compute_iou_batch(box[:4], tta_boxes[:, :4]) if (ious 0.5).any(): tta_count 1 if tta_count 3: continue # 层3几何合理性 area (x2 - x1) * (y2 - y1) ar (x2 - x1) / (y2 - y1) if (y2 - y1) 0 else 0 if area area_thresh or abs(ar - aspect_ratio_med) ar_delta: continue valid_mask[i] True return pred_boxes[valid_mask] def compute_iou_batch(box, boxes): 向量化IoU计算 x1, y1, x2, y2 box x1s, y1s, x2s, y2s boxes.T inter_x1 np.maximum(x1, x1s) inter_y1 np.maximum(y1, y1s) inter_x2 np.minimum(x2, x2s) inter_y2 np.minimum(y2, y2s) inter_area np.maximum(0, inter_x2 - inter_x1) * np.maximum(0, inter_y2 - inter_y1) box_area (x2 - x1) * (y2 - y1) boxes_area (x2s - x1s) * (y2s - y1s) union_area box_area boxes_area - inter_area return inter_area / (union_area 1e-6)4.2 伪标签修正用CRF优化分割边界检测框伪标签粗糙而分割任务需要精确边界。我们用条件随机场CRF对预测掩膜做后处理利用像素间空间与颜色相似性细化边缘import pydensecrf.densecrf as dcrf from pydensecrf.utils import unary_from_labels, create_pairwise_bilateral def crf_refine(mask_pred, image, n_iters5, compat10): mask_pred: (H, W) uint80为背景1为缺陷 image: (H, W, 3) uint8 h, w mask_pred.shape # 构建一元势能unary potential labels mask_pred.astype(np.int32) U unary_from_labels(labels, 2, gt_prob0.7) # 2类背景/缺陷 # 构建二元势能bilateral filter d dcrf.DenseCRF2D(w, h, 2) d.setUnaryEnergy(U) # 颜色空间相似性 feats create_pairwise_bilateral( sdims(80, 80), # 空间尺度 schan(13, 13, 13), # 颜色尺度RGB各通道 imgimage, chdim2 ) d.addPairwiseEnergy(feats, compatcompat) # 迭代优化 Q d.inference(n_iters) refined_mask np.argmax(Q, axis0).reshape(h, w) return refined_mask.astype(np.uint8)参数选择sdims(80,80)对应缺陷平均尺寸45px的1.8倍确保CRF能覆盖整个缺陷schan(13,13,13)来自产线图RGB通道标准差统计值R:12.3, G:11.7, B:13.1。4.3 迭代训练闭环每轮只增10%伪标签防雪崩效应伪标签引入噪声必须控制增量。我们采用渐进式注入第1轮用原始127张图训练第2轮加入12张10%高质量伪标签图第3轮再加12张同时剔除上轮中置信度下降的伪标签监控验证集mAP若下降0.5%则回滚# 伪标签管理器简化版 class PseudoLabelManager: def __init__(self, base_dataset, pseudo_dir): self.base_dataset base_dataset # 原始Dataset self.pseudo_dir pseudo_dir self.pseudo_list [] # [(img_path, ann_path, conf_score), ...] self.history [] # 每轮的pseudo_list快照 def add_batch(self, new_pseudo_list, max_add12): # 按conf_score降序取top max_add sorted_pseudo sorted(new_pseudo_list, keylambda x: x[2], reverseTrue) to_add sorted_pseudo[:max_add] # 检查是否与base_dataset重复避免同图多次加入 base_names {os.path.basename(x[0]) for x in self.base_dataset.img_files} to_add [p for p in to_add if os.path.basename(p[0]) not in base_names] self.pseudo_list.extend(to_add) self.history.append(copy.deepcopy(self.pseudo_list)) def get_current_dataset(self): # 返回当前轮次的完整数据集base pseudo return ConcatDataset([self.base_dataset, PseudoDataset(self.pseudo_list)])关键纪律每轮训练后必须用未参与训练的held-out验证集预留10张产线图评估mAP若下降则立即self.pseudo_list self.history[-2]回滚。这是防止伪标签污染的后悔药。5. 避坑指南小样本增强的五个致命翻车点小样本数据增强不是“加得越多越好”而是“加得越准越稳”。以下是我踩过的坑按发生频率排序每条都附真实复现案例5.1 现象增强后验证集mAP不升反降5%以上原因用了RandomResizedCrop且scale(0.8, 1.0)导致小缺陷平均尺寸20px被随机裁剪掉。原始127张图中32张含微小裂纹15px增强后这些图在训练中几乎不出现。解决禁用RandomResizedCrop改用CenterCrop固定裁剪尺寸原始图最小边长的0.9倍或改用Albumentations的RandomSizedBBoxSafeCrop它会确保所有标注框完整保留在裁剪区内。5.2 现象模型在测试集上召回率极高98%但精确率暴跌至42%原因MixUp增强中α0.4导致两图混合后缺陷边界模糊模型学会“只要图里有疑似区域就打框”。在产线图上这表现为把阴影、划痕、反光全标为缺陷。解决停用MixUp改用MosaicYOLO系专用增强它拼接4图但保持各图缺陷区域物理独立若必须用MixUpα必须≤0.1且只对同类缺陷图混合如锈迹图只与锈迹图mix。5.3 现象CycleGAN生成图在验证集上mAP提升2.1%但在新产线相机上掉点8.3%原因CycleGAN训练时用了--preprocess resize_and_crop把产线图统一resize到256×256破坏了原始分辨率下的缺陷纹理如0.1mm裂纹在256图中仅占2像素。解决改用--preprocess scale_width_and_crop先等比缩放至短边256再随机裁256×256保留纹理密度生成后用双三次插值缩回原始分辨率。5.4 现象伪标签筛选后得到87个框但人工检查发现23个是误检误检率26.4%原因TTA一致性过滤只用了IoU0.5但产线图中缺陷形变大如弯曲管道上的裂纹同一缺陷在不同增强下预测框IoU常0.4。解决将TTA过滤改为中心点距离过滤计算5次预测框中心点的欧氏距离标准差若15像素缺陷平均尺寸1/3则通过比IoU更鲁棒。5.5 现象CRF优化后分割掩膜边缘更锐利但检测框定位误差增大原因CRF的compat10过大过度平滑导致缺陷区域收缩检测器回归时锚点偏移。解决compat值必须与缺陷尺寸匹配——公式为compat 10 * (defect_avg_size / 50)你数据中缺陷平均尺寸45px故compat9同时CRF只对伪标签掩膜做不对原始标注做。6. 验证与调优用“缺陷敏感型评估”代替泛化指标小样本场景下mAP0.5这种通用指标会掩盖致命问题模型可能把所有缺陷都标成大框高召回低精度或只检出大缺陷漏检微小裂纹。我们必须设计缺陷特异性评估协议直接回答老板最关心的问题“它能不能在产线上不漏检”6.1 构建缺陷敏感型验证集从127张原始图中人工挑出12张“压力测试图”要求① 含微小缺陷尺寸15px② 多缺陷重叠如锈迹划痕③ 极端光照强反光/背光④ 模糊运动相机抖动。这12张不参与任何训练/增强只用于最终验证。评估时不仅看mAP更盯三个硬指标指标计算方式合格线业务意义微缺陷召回率尺寸15px的缺陷中被正确检出的比例≥85%决定是否需返工重叠缺陷分离度两个中心距30px的缺陷被分框检出的比例≥70%避免误判为单一大缺陷强反光鲁棒性反光区域HSV V220且S30内缺陷的检出率≥90%产线常见干扰def evaluate_defect_sensitivity(model, val_loader, stress_test_images): stress_test_images: List of paths to 12 pressure-test images 返回字典{micro_recall: 0.87, overlap_separation: 0.72, glare_robust: 0.91} results {micro_recall: [], overlap_separation: [], glare_robust: []} for img_path in stress_test_images: image cv2.imread(img_path) # 获取原始标注人工精标 gt_ann load_gt_annotation(img_path) # 返回List[{bbox: [x1,y1,x2,y2], size: px, type: rust}] # 模型预测 preds model.predict(image) # 微缺陷召回gt中size15px的缺陷pred中IoU0.5的框数/总数 micro_gt [g for g in gt_ann if g[size] 15] micro_hit 0 for g in micro_gt: ious [compute_iou(g[bbox], p[:4]) for p in preds] if any(i 0.5 for i in ious): micro_hit 1 results[micro_recall].append(micro_hit / len(micro_gt) if micro_gt else 1.0) # 重叠缺陷分离gt中中心距30px的缺陷对pred中分框检出比例 overlap_pairs [] for i in range(len(gt_ann)): for j in range(i1, len(gt_ann)): c1 ((gt_ann[i][bbox][0]gt_ann[i][bbox][2])/2, (gt_ann[i][bbox][1]gt_ann[i][bbox][3])/2) c2 ((gt_ann[j][bbox][0]gt_ann[j][bbox][2])/2, (gt_ann[j][bbox][1]gt_ann[j][bbox][3])/2) dist np.sqrt((c1[0]-c2[0])**2 (c1[1]-c2[1])**2) if dist 30: overlap_pairs.append((i,j)) sep_count 0 for i,j in overlap_pairs: # 检查pred中是否有两个框分别匹配i和j match_i any(compute_iou(gt_ann[i][bbox], p[:4]) 0.5 for p in preds) match_j any(compute_iou(gt_ann[j][bbox], p[:4]) 0.5 for p in preds) if match_i and match_j: sep_count 1 results[overlap_separation].append(sep_count / len(overlap_pairs) if overlap_pairs else 1.0) # 强反光鲁棒性gt缺陷在反光区内的检出率 glare_mask create_glare_mask(image) # HSV V220 S30 glare_gt [g for g in gt_ann if is_in_mask(g[bbox], glare_mask)] glare_hit sum(1 for g in glare_gt if any(compute_iou(g[bbox], p[:4]) 0.5 for p in preds)) results[glare_robust].append(glare_hit / len(glare_gt) if glare_gt else 1.0) # 返回均值 return {k: np.mean(v) for k,v in results.items()}6.2 增强效果归因分析用Grad-CAM定位“模型到底学到了什么”增强是否有效不能只看指标要看模型注意力是否聚焦在真实缺陷上。我们用Grad-CAM可视化最后一层卷积的梯度响应def gradcam_visualize(model, image, target_layermodel.model[10]): # YOLOv8的neck层 image: (3, H, W) tensor, 归一化 返回热力图H, W值越大表示该区域对预测贡献越大 model.eval() image.requires_grad_(True) # 获取目标层输出 target_activations None def hook_fn(module, input, output): nonlocal target_activations target_activations output hook eval(fmodel.{target_layer}).register_forward_hook(hook_fn) # 前向传播 preds model(image.unsqueeze(0)) # 取最高置信度预测的类别得分 scores preds[0][:, 4] * preds[0][:, 5:].max(dim1)[0] # conf * cls_score if len(scores) 0: return np.zeros(image.shape[1:]) top_score_idx scores.argmax().item() top_score scores[top_score_idx] # 反向传播 p a hrefhttps://download.csdn.net/download/weixin_38693528/13752142 stylecolor:#ec7500;font-size:14px; 本文还有配套的精品资源点击获取 /a img altmenu-r.4af5f7ec.gif srchttps://csdnimg.cn/release/wenkucmsfe/public/img/menu-r.4af5f7ec.gif stylewidth:16px;margin-left:4px;vertical-align:text-bottom;cursor:text; /p
返回列表