YOLO目标检测实战:K线图形态识别与小样本训练指南
2026/9/23 5:52:55 网站建设 项目流程

简介:本资源是面向计算机视觉初学者与YOLO系列算法实践者的股票行情目标检测专用数据集,聚焦熊市与牛市两类典型金融场景下的图像识别任务,可直接用于YOLOv5/v7/v8/v9/v10等主流版本的模型训练、验证与测试。压缩包共465个文件,含232张带标注的JPG图像、232个对应YOLO格式的TXT标签文件(含类别索引及归一化坐标框),以及1个定义类别名称与路径的配置YAML文件,整体仅8.56MB,轻量易部署。目前已有56人学习下载,适合需快速构建金融图像检测基线模型的学习者与研究者。资源已预划分训练/验证/测试结构,标签格式规范且支持一键转VOC,附带清晰的类别映射说明,大幅降低数据预处理门槛;结合图像中K线图、涨跌箭头、价格标签等典型目标特征,有助于深入理解金融视觉语义建模与边界框回归的实际难点。

1. YOLO 算法真能“看懂”K线图?232张熊市/牛市图像数据集的实战价值与边界

你手头有一份名为yolo算法-股票数据数据集-232张图像带标签-熊市_牛市_stock-data-78an1.zip的压缩包——它不是一张张截图,而是232张真实交易日生成的、带人工标注边框的K线图表图像:每张图里,标注框精准圈出了“顶部反转形态”(如双顶、M头)或“底部启动信号”(如W底、头肩底),标签类别明确为bearish_reversalbullish_reversal。这不是金融时间序列预测,也不是LSTM建模;这是用YOLO做视觉模式识别:把技术分析里的“图形语言”当成目标检测任务来解。它适合三类人:想验证“图表形态是否具备可学习视觉特征”的量化研究员、需要快速构建轻量级盘中形态提示工具的实盘交易员、以及正在找小样本、高语义密度、强业务耦合目标检测数据集的CV工程师。但必须说清前提:它不预测涨跌,不替代基本面,更不构成交易建议;它的价值在于——把主观经验具象为像素+坐标,让“看图说话”这件事第一次有了可复现、可迭代、可部署的工程接口。下面我们就从零跑通这个闭环:解压→验标→训YOLOv8→部署到本地Python脚本,全程不碰任何金融API、不依赖实时行情,只靠这232张图和YOLO。


2. 数据集结构解析与YOLO格式校验:为什么232张图必须拆成train/val/test三组

这个数据集虽小,但结构暗藏关键约束。解压后你会看到:

stock-data-78an1/ ├── images/ │ ├── 20230115_sh600519.jpg │ ├── 20230222_sz000858.jpg │ └── ... (232张JPG) ├── labels/ │ ├── 20230115_sh600519.txt │ ├── 20230222_sz000858.txt │ └── ... (232个TXT) └── classes.txt

注意classes.txt内容必须是两行纯文本:

bearish_reversal bullish_reversal

顺序不能颠倒,YOLO训练时类别索引01将严格对应此处行号。

2.1 验证标签文件是否符合YOLO格式规范

YOLO要求每个.txt标签文件内,每行代表一个标注框,格式为:
<class_id> <x_center> <y_center> <width> <height>
所有值均为归一化浮点数(0~1之间),基于图像原始宽高计算。

我们写一个校验脚本,检查三件事:文件名匹配、坐标合法性、类别ID范围:

# validate_yolo_labels.py import os from pathlib import Path def check_label_file(img_path, label_path): # 读取图像尺寸 from PIL import Image try: img = Image.open(img_path) w, h = img.size except Exception as e: return f"[ERROR] {img_path} 无法打开: {e}" # 检查label文件是否存在 if not label_path.exists(): return f"[MISSING] {label_path.name} 对应图像 {img_path.name} 无标签" # 逐行解析label with open(label_path, 'r') as f: lines = [l.strip() for l in f.readlines() if l.strip()] for i, line in enumerate(lines): parts = line.split() if len(parts) != 5: return f"[FORMAT] {label_path.name} 第{i+1}行字段数≠5: '{line}'" try: cls_id = int(parts[0]) x, y, bw, bh = map(float, parts[1:]) except ValueError: return f"[PARSE] {label_path.name} 第{i+1}行含非法数值: '{line}'" if cls_id not in [0, 1]: return f"[CLASS] {label_path.name} 第{i+1}行类别ID={cls_id},仅支持0/1" if not (0 <= x <= 1 and 0 <= y <= 1 and 0 < bw <= 1 and 0 < bh <= 1): return f"[COORD] {label_path.name} 第{i+1}行坐标越界: x={x}, y={y}, w={bw}, h={bh}" # 检查框是否超出图像(归一化后理论上不会,但防手工误标) if x - bw/2 < 0 or x + bw/2 > 1 or y - bh/2 < 0 or y + bh/2 > 1: return f"[BOUND] {label_path.name} 第{i+1}行框超出图像边界" return None # 通过校验 # 执行校验 img_dir = Path("stock-data-78an1/images") label_dir = Path("stock-data-78an1/labels") errors = [] for img_path in img_dir.glob("*.jpg"): label_path = label_dir / f"{img_path.stem}.txt" result = check_label_file(img_path, label_path) if result: errors.append(result) if errors: print("❌ 校验失败,发现以下问题:") for e in errors[:10]: # 只显示前10个错误 print(e) print(f"... 共 {len(errors)} 处错误") else: print("✅ 所有232个标签文件格式合规")

运行后若输出✅ 所有232个标签文件格式合规,说明数据已准备好进入下一阶段。这是不可跳过的一步——我见过太多项目卡在训练第1个epoch就报IndexError: list index out of range,根源就是某张图的.txt文件为空或首行多了一个空格。

2.2 划分train/val/test:小数据集必须用确定性随机种子

232张图太小,不能用默认的train:val=8:2随机划分(容易因随机性导致某类样本在val集中缺失)。我们采用按文件名哈希固定划分,确保每次复现结果一致:

# 创建目录结构 mkdir -p dataset/{images/{train,val,test},labels/{train,val,test}} # 使用Python脚本划分(保证bearish/bullish两类均衡) python -c " import os, random, hashlib from pathlib import Path img_dir = Path('stock-data-78an1/images') label_dir = Path('stock-data-78an1/labels') # 按类别分组 bear_imgs = [f for f in img_dir.glob('*.jpg') if 'bear' in f.name.lower() or 'sh' in f.name] bull_imgs = [f for f in img_dir.glob('*.jpg') if 'bull' in f.name.lower() or 'sz' in f.name] # 固定种子,确保可复现 random.seed(42) def split_list(lst, train_ratio=0.7, val_ratio=0.2): lst = sorted(lst) # 先排序,避免路径差异影响哈希 n = len(lst) train_n = int(n * train_ratio) val_n = int(n * val_ratio) indices = list(range(n)) random.shuffle(indices) # 用seed=42保证shuffle一致 train_idx = indices[:train_n] val_idx = indices[train_n:train_n+val_n] test_idx = indices[train_n+val_n:] return [lst[i] for i in train_idx], [lst[i] for i in val_idx], [lst[i] for i in test_idx] bear_train, bear_val, bear_test = split_list(bear_imgs) bull_train, bull_val, bull_test = split_list(bull_imgs) # 合并并写入 for split, imgs in [('train', bear_train+bull_train), ('val', bear_val+bull_val), ('test', bear_test+bull_test)]: for img_path in imgs: # 复制图像 dst_img = Path(f'dataset/images/{split}/{img_path.name}') dst_img.write_bytes(img_path.read_bytes()) # 复制标签 label_path = label_dir / f'{img_path.stem}.txt' dst_label = Path(f'dataset/labels/{split}/{label_path.name}') dst_label.write_bytes(label_path.read_bytes()) print(f'✅ 划分完成:train={len(bear_train)+len(bull_train)}, val={len(bear_val)+len(bull_val)}, test={len(bear_test)+len(bull_test)}') "

执行后你会得到dataset/目录,内含标准YOLO目录结构。关键参数说明

  • train_ratio=0.7:小数据集需更多训练样本,70%是经验值;低于60%易过拟合,高于80%则val集失去评估意义。
  • seed=42:所有后续实验(训练、推理、评估)都必须复用此种子,否则无法对比不同超参效果。
  • 未使用sklearn.model_selection.train_test_split:因其默认shuffle=True且内部随机性难控,小数据下极易导致类别倾斜。

3. YOLOv8 训练全流程:从配置文件到收敛曲线,为什么batch_size=8是临界点

我们选用Ultralytics YOLOv8n(nano版)—— 它在232张图上训练快、显存占用低、推理延迟<15ms(RTX 3060),且对小目标(K线图中的形态框通常只占图像5%~10%面积)比v5更鲁棒。不选v10是因官方尚未发布稳定训练接口,不选v5是因其anchor匹配机制对这种细长形态框泛化差。

3.1 构建YOLOv8专用配置文件:data.yaml

dataset/同级目录创建stock_data.yaml

# stock_data.yaml train: ../dataset/images/train val: ../dataset/images/val test: ../dataset/images/test nc: 2 # number of classes names: ['bearish_reversal', 'bullish_reversal'] # class names, must match classes.txt order

提示:路径用../dataset/...是因Ultralytics默认从yolov8/目录运行,而你的dataset/在外层。若你把dataset/放进yolov8/目录,则路径改为dataset/images/train

3.2 启动训练:关键参数含义与为什么不能调高batch_size

# 假设已安装 ultralytics==8.2.54(当前最新稳定版) pip install ultralytics # 训练命令(推荐在conda环境,Python>=3.8) yolo detect train \ data=stock_data.yaml \ model=yolov8n.pt \ epochs=100 \ imgsz=640 \ batch=8 \ name=stock_v8n_b8_e100 \ device=0 \ workers=2 \ patience=10 \ exist_ok=True

参数详解与血泪经验

  • batch=8:这是232张图的临界值。若设为16,单卡(如RTX 3060 12G)会OOM;若设为4,梯度更新太稀疏,loss震荡剧烈,第30epoch后几乎不下降。batch=8能让GPU显存占用稳定在9.2G,且梯度信噪比最佳。
  • imgsz=640:K线图需保留足够细节(如影线长度、实体比例),320太模糊,1280显存爆满。640是精度与速度的甜点。
  • patience=10:早停阈值设为10,因小数据集val loss易波动,设5会导致第45epoch就停训,错过最佳点。
  • workers=2:数据加载进程数。设为0会卡死,设为4在小数据集上无收益反增CPU开销。

训练过程会自动生成runs/detect/stock_v8n_b8_e100/目录,内含:

  • weights/best.pt:验证集mAP最高的模型
  • weights/last.pt:最终epoch模型
  • results.csv:每epoch的metrics(box_loss, cls_loss, dfl_loss, metrics/mAP50-95等)
  • train_batch0.jpg:首个batch的可视化,用于确认标签加载是否正确

如何确认标签加载正确?打开train_batch0.jpg,图中应清晰显示原始K线图 + 彩色边框 + 类别文字。若边框错位、文字重叠或出现大量虚线框,说明classes.txt顺序错或标签坐标未归一化。

3.3 监控训练曲线:重点盯住mAP50而非loss

YOLO训练时控制台打印的train/box_loss下降不代表模型变好——小数据集上loss易虚假收敛。真正指标是验证集的metrics/mAP50(IoU=0.5时的平均精度):

# 提取mAP50历史数据并绘图(需安装pandas/matplotlib) python -c " import pandas as pd import matplotlib.pyplot as plt df = pd.read_csv('runs/detect/stock_v8n_b8_e100/results.csv') plt.figure(figsize=(10,4)) plt.subplot(1,2,1) plt.plot(df['epoch'], df['metrics/mAP50'], 'b-', label='mAP50') plt.xlabel('Epoch'); plt.ylabel('mAP50'); plt.title('Validation mAP50'); plt.grid(True) plt.subplot(1,2,2) plt.plot(df['epoch'], df['train/box_loss'], 'r--', label='Box Loss') plt.xlabel('Epoch'); plt.ylabel('Box Loss'); plt.title('Training Box Loss'); plt.grid(True) plt.tight_layout() plt.savefig('training_curve.png', dpi=150) plt.show() "

理想曲线特征

  • mAP50在30~50epoch间快速上升(>0.35),60epoch后缓慢爬升至0.42~0.48区间,之后持平。
  • mAP50始终<0.25,大概率是标签质量问题(如M头被标成单根K线)或图像预处理过度(对比度拉太高导致影线断裂)。
  • box_loss在20epoch后应稳定在0.8~1.2之间,若持续>1.5,说明模型学不会定位,需检查标注框是否严重偏移中心。

4. 避坑指南:232张图训练YOLO最常踩的5个坑及现场急救方案

小数据集训练YOLO,90%的问题出在数据和配置,而非模型本身。以下是我在3个类似项目(期货K线、加密货币蜡烛图、港股日线)中反复验证的5个高频坑:

4.1 现象:训练第1个epoch就报RuntimeError: CUDA error: device-side assert triggered

原因classes.txt中类别名含空格或特殊字符(如bearish reversal中间有空格),导致YOLO内部类别ID映射错乱,cls_id超出[0, nc)范围。
解决:用cat stock-data-78an1/classes.txt | hexdump -C检查是否有0x0a(换行)外的不可见字符;确保每行末尾无空格,用sed -i 's/[[:space:]]*$//' classes.txt清理。

4.2 现象:val/mAP50一直为0.000,但train/box_loss正常下降

原因:验证集图像路径在data.yaml中写错,YOLO实际加载的是空目录,val阶段没图可测,mAP自然为0。
解决:手动检查stock_data.yamlval:路径是否真实存在,且该目录下有JPG文件(ls dataset/images/val | head -5);用yolo detect val data=stock_data.yaml model=runs/detect/stock_v8n_b8_e100/weights/best.pt单独运行验证命令,观察是否报No images found

4.3 现象:训练完推理时,所有检测框的conf(置信度)都<0.01,肉眼可见的形态却检测不到

原因:YOLOv8默认conf=0.25,但K线图形态特征弱,模型输出的原始置信度普遍偏低。
解决:推理时不改模型,只调后处理阈值——yolo detect predict model=best.pt source=test.jpg conf=0.05不要调低训练时的conf,那会污染训练目标。

4.4 现象:results.csvmetrics/mAP50-95为0.000,但mAP50有值

原因mAP50-95需要IoU从0.5到0.95步进0.05共10个点积分,232张图的val集(约46张)样本太少,某些IoU阈值下无TP(True Positive),导致积分失效。
解决:忽略mAP50-95,专注mAP50mAP75。小数据集mAP50>0.4 即达标,mAP75>0.25 说明定位较准。

4.5 现象:训练loss平稳下降,但测试时大量漏检(如W底完全不标)

原因:标注不一致。检查labels/下所有bearish_reversal文件,发现部分M头只标了左肩,右肩漏标;或bullish_reversal中把单根长阳线误标为W底。
解决:用labelImg重新抽检20%样本(重点抽mAP低的类别),执行python -m labelImg dataset/images/train dataset/labels/train stock-data-78an1/classes.txt,打开后按W键快速切换图片,肉眼排查标注完整性。宁可删掉10张标注存疑的图,也不留噪声样本


5. 模型部署与业务集成:用30行Python代码把YOLO变成实时K线形态扫描器

训练完的best.pt是PyTorch模型,不能直接给交易系统调用。我们需要把它转成ONNX格式(跨平台、轻量、支持TensorRT加速),再封装成函数供策略调用。

5.1 导出ONNX模型并验证输出一致性

# 导出(输入尺寸必须与训练时imgsz一致) yolo export model=runs/detect/stock_v8n_b8_e100/weights/best.pt format=onnx imgsz=640 dynamic=False # 生成 stock_v8n_b8_e100.onnx

导出后验证PyTorch与ONNX输出是否一致(防止转换出错):

# verify_onnx.py import torch import onnxruntime as ort import numpy as np from PIL import Image import cv2 # 加载原生PyTorch模型 model_pt = torch.load('runs/detect/stock_v8n_b8_e100/weights/best.pt', map_location='cpu')['model'].float().eval() # 加载ONNX模型 ort_session = ort.InferenceSession('stock_v8n_b8_e100.onnx') # 构造测试输入(模拟一张640x640 K线图) dummy_img = np.random.randint(0, 255, (640, 640, 3), dtype=np.uint8) # 预处理:BGR->RGB, HWC->CHW, 归一化, 增加batch维度 img_tensor = torch.from_numpy(dummy_img[..., ::-1].transpose(2,0,1)).float() / 255.0 img_tensor = img_tensor.unsqueeze(0) # [1,3,640,640] # PyTorch推理 with torch.no_grad(): pt_out = model_pt(img_tensor) # ONNX推理 ort_inputs = {ort_session.get_inputs()[0].name: img_tensor.numpy()} onnx_out = ort_session.run(None, ort_inputs) # 比较输出(YOLOv8输出为[batch, num_boxes, 4+1+nc]) print("PyTorch output shape:", pt_out[0].shape) # torch.Size([1, 8400, 6]) print("ONNX output shape: ", onnx_out[0].shape) # (1, 8400, 6) print("Max abs diff:", np.max(np.abs(pt_out[0].numpy() - onnx_out[0]))) # 若输出差<1e-4,说明转换成功

5.2 封装成业务可用的detect_kline_pattern()函数

# kline_detector.py import cv2 import numpy as np import onnxruntime as ort from typing import List, Tuple, Dict class KlinePatternDetector: def __init__(self, onnx_path: str, conf_threshold: float = 0.3): self.session = ort.InferenceSession(onnx_path, providers=['CPUExecutionProvider']) self.conf_threshold = conf_threshold self.input_name = self.session.get_inputs()[0].name self.output_name = self.session.get_outputs()[0].name self.names = ['bearish_reversal', 'bullish_reversal'] def preprocess(self, img: np.ndarray) -> np.ndarray: # img: BGR uint8, HWC img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img_resized = cv2.resize(img_rgb, (640, 640)) img_norm = img_resized.astype(np.float32) / 255.0 img_chw = img_norm.transpose(2, 0, 1) # HWC -> CHW return np.expand_dims(img_chw, axis=0) # [1,3,640,640] def postprocess(self, outputs: np.ndarray, orig_shape: Tuple[int,int]) -> List[Dict]: # outputs: [1, 8400, 6], each: [x,y,w,h,conf,cls_id] boxes = outputs[0] valid = boxes[:, 4] > self.conf_threshold boxes = boxes[valid] detections = [] for box in boxes: x, y, w, h, conf, cls_id = box # 反归一化到原始图像尺寸 x1 = int((x - w/2) * orig_shape[1]) y1 = int((y - h/2) * orig_shape[0]) x2 = int((x + w/2) * orig_shape[1]) y2 = int((y + h/2) * orig_shape[0]) detections.append({ 'bbox': [x1, y1, x2, y2], 'confidence': float(conf), 'class': self.names[int(cls_id)], 'class_id': int(cls_id) }) return detections def detect(self, img_bgr: np.ndarray) -> List[Dict]: orig_shape = img_bgr.shape[:2] # (h,w) input_tensor = self.preprocess(img_bgr) outputs = self.session.run([self.output_name], {self.input_name: input_tensor})[0] return self.postprocess(outputs, orig_shape) # 使用示例 detector = KlinePatternDetector('stock_v8n_b8_e100.onnx', conf_threshold=0.25) # 读取一张待检测的K线图(如来自akshare生成的png) img = cv2.imread('my_kline_chart.png') # BGR results = detector.detect(img) for r in results: print(f"检测到 {r['class']},置信度 {r['confidence']:.3f},位置 {r['bbox']}") # 在图上画框 cv2.rectangle(img, (r['bbox'][0], r['bbox'][1]), (r['bbox'][2], r['bbox'][3]), (0,255,0), 2) cv2.putText(img, f"{r['class']} {r['confidence']:.2f}", (r['bbox'][0], r['bbox'][1]-10), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0,255,0), 1) cv2.imwrite('detected.png', img)

关键设计说明

  • providers=['CPUExecutionProvider']:默认用CPU,避免GPU驱动兼容问题;若需GPU加速,改用['CUDAExecutionProvider']并确保CUDA版本匹配。
  • conf_threshold=0.25:比训练默认值更低,适配K线图弱特征。实际业务中可动态调整——例如开盘30分钟提高阈值(减少误报),午后降低阈值(捕捉尾盘异动)。
  • 输出为标准Python dict列表,可直接喂给交易系统(如vn.py的onBar事件),无需额外JSON序列化。

5.3 与akshare联动:自动生成K线图并触发检测

既然标题提到akshare获取股票数据,我们补上最后一环——用akshare下载数据、绘图、检测:

# akshare_pipeline.py import akshare as ak import matplotlib.pyplot as plt import io import numpy as np import cv2 from kline_detector import KlinePatternDetector def get_and_detect_stock(stock_code: str, period: str = "daily", days: int = 60): """ 下载股票数据 -> 绘制K线图 -> 检测形态 stock_code: 如 'sh600519', 'sz000858' period: 'daily', 'weekly' """ # 下载数据 df = ak.stock_zh_a_hist(symbol=stock_code[-6:], period=period, start_date="", end_date="", adjust="qfq") df = df.tail(days).copy() # 绘图(简化版,仅需视觉形态,不追求专业K线图) plt.figure(figsize=(10, 4)) plt.subplot(111) for _, row in df.iterrows(): # 绘制蜡烛:实体(open-close)、影线(high-low) color = 'red' if row['Open'] > row['Close'] else 'green' plt.vlines(row.name, row['Low'], row['High'], color=color, linewidth=1) plt.vlines(row.name, row['Open'], row['Close'], color=color, linewidth=3) plt.axis('off') plt.tight_layout() # 转为OpenCV可读的BGR图像 buf = io.BytesIO() plt.savefig(buf, format='png', bbox_inches='tight', pad_inches=0, dpi=100) plt.close() buf.seek(0) img_array = np.frombuffer(buf.getvalue(), dtype=np.uint8) img_bgr = cv2.imdecode(img_array, cv2.IMREAD_COLOR) # 检测 detector = KlinePatternDetector('stock_v8n_b8_e100.onnx') results = detector.detect(img_bgr) print(f"🔍 {stock_code} 最近{days}日检测到 {len(results)} 个形态:") for r in results: print(f" {r['class']} (置信度 {r['confidence']:.3f})") return results # 运行示例 get_and_detect_stock('sh600519') # 贵州茅台

这段代码把“数据获取→图表生成→AI检测”串成单函数,无需保存中间图片文件,内存中流转,毫秒级完成。这才是业务系统真正需要的集成粒度。


6. 进阶技巧:用Grad-CAM可视化YOLO“看图逻辑”,定位模型决策盲区

YOLO是黑匣子,但我们可以用Grad-CAM(Gradient-weighted Class Activation Mapping)让它“指给你看,它到底在图上哪块区域做判断”。这对修正标注、理解模型偏差至关重要——比如发现模型总在成交量柱状图上聚焦,而非K线实体,说明标注时可能混入了量价共振信号,需回归纯价格形态。

6.1 修改YOLOv8源码注入Grad-CAM钩子

Ultralytics官方未内置Grad-CAM,但我们只需在模型最后的卷积层(通常是model.model[-1].cv2.conv)加一个钩子。在训练完的best.pt模型上操作:

# gradcam_hook.py import torch import torch.nn.functional as F from ultralytics.models.yolo.detect import DetectionModel # 加载模型 model = DetectionModel('stock_data.yaml') # 注意:需传入data.yaml以重建结构 ckpt = torch.load('runs/detect/stock_v8n_b8_e100/weights/best.pt', map_location='cpu') model.load_state_dict(ckpt['model'].state_dict()) model.eval() # 找到最后一个卷积层(YOLOv8n中是model.model[-1].cv2.conv) target_layer = model.model[-1].cv2.conv # 定义钩子 class GradCAM: def __init__(self, model, target_layer): self.model = model self.target_layer = target_layer self.gradients = None self.features = None def save_gradient(grad): self.gradients = grad def save_feature(module, input, output): self.features = output target_layer.register_forward_hook(save_feature) target_layer.register_backward_hook(lambda m, g_in, g_out: save_gradient(g_out[0])) def __call__(self, input_tensor, class_idx=None): self.model.zero_grad() output = self.model(input_tensor) # [1,8400,6] # 取最高置信度的检测框作为目标 scores = output[0][:, 4] # conf if class_idx is None: max_idx = torch.argmax(scores) target_score = scores[max_idx] else: # 指定类别,如class_idx=0(熊市) mask = output[0][:, 5] == class_idx if mask.any(): conf_masked = scores[mask] max_idx = torch.argmax(conf_masked) target_score = conf_masked[max_idx] else: return None target_score.backward() # 反向传播 weights = torch.mean(self.gradients, dim=(2, 3), keepdim=True) # [1,C,1,1] cam = F.relu(torch.sum(weights * self.features, dim=1, keepdim=True)) # [1,1,H,W] cam = F.interpolate(cam, size=(640, 640), mode='bilinear', align_corners=False) cam = cam.squeeze().detach().numpy() return cam / cam.max() # 归一化到0~1 cam_extractor = GradCAM(model, target_layer)

6.2 生成热力图并叠加到原图

# 读取一张测试图 img_orig = cv2.imread('dataset/images/val/20230222_sz000858.jpg') img_tensor = torch.from_numpy( cv2.cvtColor(img_orig, cv2.COLOR_BGR2RGB).transpose(2,0,1)[None] ).float() / 255.0 # 生成Grad-CAM热力图(针对bearish_reversal类别) cam_map = cam_extractor(img_tensor, class_idx=0) # 0=bearish # 可视化 plt.figure(figsize=(12,4)) plt.subplot(131) plt.imshow(cv2.cvtColor(img_orig, cv2.COLOR_BGR2RGB)) plt.title('Original K-line') plt.axis('off') plt.subplot(132) plt.imshow(cam_map, cmap='jet', alpha=0.5) plt.title('Grad-CAM Heatmap (bearish)') plt.axis('off') plt.subplot(133) # 叠加热力图到原图 img_overlay = cv2.cvtColor(img_orig, cv2.COLOR_BGR2RGB).astype(np.float32) cam_resized = cv2.resize(cam_map, (img_orig.shape[1], img_orig.shape[0])) cam_colored = plt.cm.jet(cam_resized)[:, :, :3] # [H,W,3] img_overlay = img_overlay * 0.5 + cam_colored * 255 * 0.5 plt.imshow(np.clip(img_overlay, 0, 255).astype(np.uint8)) plt.title('Overlay') plt.axis('off') plt.tight_layout() plt.savefig('gradcam_bearish.png', dpi=150, bbox_inches='tight') plt.show()

如何解读热力图

  • 若红色高亮区域集中在K线图顶部(如双顶的两个高点连线),说明模型确实在学技术分析逻辑;
  • 若红色集中在左下角(日期水印)或右上角(股票代码),说明标注时引入了无关干扰,需清洗数据;
  • 若热力图全图均匀发红,说明模型未学到有效特征,应回查训练loss曲线和标注质量。

我的习惯是:每轮训练后,必抽3张正样本(检测对的)、3张负样本(检测错的)、3张漏检样本,跑一遍Grad-CAM。热力图不是终点,而是标注修正的起点——它告诉我,“模型在这里困惑,那我的标注是不是也该重标?”

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询