news 2026/9/15 3:19:53

Swin-Transformer中文数据集构建与训练实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Swin-Transformer中文数据集构建与训练实战

简介:本资源是一套面向深度学习初学者与计算机视觉实践者的Swin-Transformer图像识别完整项目,覆盖从关键词驱动的网络图像采集、数据清洗与集划分,到模型训练、推理部署的全流程。项目以漫威角色(钢铁侠、美国队长、雷神)为实际案例,含347张训练图与85张测试图,实测精度达91%,配套脚本自动完成数据格式转换、类别JSON生成及预测结果输出,显著降低Transformer模型落地门槛。压缩包共475个文件,主体为390张jpeg及28张png/webp格式图像,辅以14个核心Python训练与推理脚本、2个预训练.pth模型、1个README说明文档及UI界面文件,整体875.29MB,结构清晰、开箱即用。目前已有405人学习下载,读者可直接复现端到端流程,掌握自定义数据集构建、Swin模型微调及中文标签输出等关键能力。

1. Swin-Transformer 不是“调个库就行”的图像识别:它真正吃的是你亲手筛过的中文关键词数据集

很多人以为 Swin-Transformer 是个“换 backbone 就能涨点”的黑箱模型——把 ResNet 换成 Swin-T,改两行 config,跑完就发报告。但实际落地时,90% 的精度波动不出在模型结构,而出在数据集生成环节的三个隐性断点:关键词检索结果混入无关图、图像损坏未剔除导致 DataLoader 报错中断、训练/测试集划分后标签路径与 JSON 类别映射不一致。本项目用钢铁侠、美国队长、雷神三类中文关键词(非英文 label)构建 347+85 样本数据集,全程支持中文路径、中文类别名、中文日志输出,验证了 Swin-Transformer 在小样本中文语义场景下的鲁棒性(测试精度 0.91)。它适合两类人:一是需要快速验证 Swin 架构在自有业务图(如工业零件、医疗胶片、中文文档截图)上效果的工程师;二是正卡在“下载→清洗→格式化→训练”链路某一步、反复报OSError: broken dataKeyError: 'iron_man'的初学者。所有脚本均适配 Windows/Linux/macOS,无需修改路径分隔符。

2. 从中文关键词到可用数据集:下载、校验、划分三步不可跳过

Swin-Transformer 对输入数据的结构敏感度远高于 CNN:它依赖 patch embedding 的局部一致性,一张损坏的 JPEG(如截断头、EXIF 元数据异常)会导致整个 batch 的 attention map 崩溃。本项目用纯 Python 脚本完成端到端数据准备,不依赖 Selenium 或浏览器自动化,规避反爬封 IP 风险。

2.1 中文关键词驱动的图像批量下载:绕过百度图片接口限制

百度图片搜索页(如https://image.baidu.com/search/index?tn=baiduimage&word=钢铁侠)的缩略图 URL 实际指向 CDN 地址,需解析 HTML 中的># download_images.py import requests, os, time, re from bs4 import BeautifulSoup def get_image_urls(keyword: str, max_count: int = 100) -> list: headers = { "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36" } # 百度图片搜索 URL 编码需处理中文 encoded_keyword = keyword.encode('utf-8').hex() url = f"https://image.baidu.com/search/acjson?tn=resultjson_com&ipn=rj&ct=201326592&fp=result&queryWord={keyword}&word={keyword}&pn=0&rn={max_count}" try: resp = requests.get(url, headers=headers, timeout=10) resp.raise_for_status() data = resp.json() urls = [item['thumbURL'] for item in data.get('data', []) if 'thumbURL' in item] return urls[:max_count] except Exception as e: print(f"[ERROR] 下载 {keyword} 失败: {e}") return [] # 批量下载三类角色 keywords = ["钢铁侠", "美国队长", "雷神"] for kw in keywords: urls = get_image_urls(kw, max_count=120) # 每类预留冗余 save_dir = os.path.join("raw_images", kw) os.makedirs(save_dir, exist_ok=True) for i, url in enumerate(urls): try: img_data = requests.get(url, timeout=5).content ext = ".jpg" if b"JFIF" in img_data[:10] else ".png" with open(os.path.join(save_dir, f"{kw}_{i:04d}{ext}"), "wb") as f: f.write(img_data) time.sleep(0.3) # 降低请求频率 except Exception as e: print(f"[SKIP] {url} 下载失败: {e}")

提示:百度接口返回的thumbURL是缩略图,项目后续通过PIL.Image.open().convert('RGB')自动转为标准 RGB 格式,避免 PNG 透明通道干扰 Swin 的 patch 切分。若需高清原图,可将thumbURL替换为objURL(但成功率下降约 40%,需加 try-except 降级)。

2.2 图像完整性校验与自动修复:剔除损坏文件并统一尺寸

validate_and_resize.py脚本执行三项关键操作:

  1. 损坏检测:用PIL.Image.open()尝试加载,捕获OSError(如 truncated image)、IOError(如 invalid JPEG data);
  2. 尺寸归一化:Swin-Transformer 默认输入为 224×224,但原始图宽高比差异大,直接 resize 会拉伸失真。本项目采用center-crop + pad策略:先按短边缩放至 256,再中心裁剪 224×224,不足部分用均值填充(RGB 均值取 [123.675, 116.28, 103.53]);
  3. 中文路径兼容:显式指定encoding='utf-8',避免 Windows 下os.listdir()返回乱码。
# validate_and_resize.py from PIL import Image, ImageOps import numpy as np import os, glob def safe_load_image(path: str) -> Image.Image: try: img = Image.open(path).convert('RGB') img.verify() # 触发 EXIF 解析,暴露隐藏损坏 return img except Exception as e: print(f"[CORRUPT] {path} 已损坏,已跳过") return None def resize_with_pad(img: Image.Image, target_size=(224, 224), fill_color=(123, 116, 103)) -> Image.Image: # 按短边缩放到 256,保持宽高比 w, h = img.size scale = 256 / min(w, h) new_w, new_h = int(w * scale), int(h * scale) img = img.resize((new_w, new_h), Image.BICUBIC) # 中心裁剪 224x224 left = (new_w - target_size[0]) // 2 top = (new_h - target_size[1]) // 2 right = left + target_size[0] bottom = top + target_size[1] img = img.crop((left, top, right, bottom)) return img # 主流程 raw_root = "raw_images" clean_root = "dataset" os.makedirs(clean_root, exist_ok=True) for class_name in os.listdir(raw_root): class_path = os.path.join(raw_root, class_name) if not os.path.isdir(class_path): continue # 创建 clean 子目录 clean_class = os.path.join(clean_root, class_name) os.makedirs(clean_class, exist_ok=True) for img_path in glob.glob(os.path.join(class_path, "*.*")): img = safe_load_image(img_path) if img is None: continue try: resized = resize_with_pad(img) # 保存为 JPG,强制统一格式(避免 PNG alpha 通道) save_name = os.path.basename(img_path).rsplit('.', 1)[0] + ".jpg" resized.save(os.path.join(clean_class, save_name), "JPEG", quality=95) except Exception as e: print(f"[RESIZE_FAIL] {img_path} 处理失败: {e}")

2.3 训练/测试集划分与目录结构生成:自动生成 Swin 兼容的 ImageFolder 格式

Swin-Transformer 官方实现(如 PyTorch Image Models)默认使用torchvision.datasets.ImageFolder,要求目录结构为:

dataset/ ├── 钢铁侠/ │ ├── 0001.jpg │ └── ... ├── 美国队长/ │ └── ... └── 雷神/ └── ...

ImageFolder无法直接指定 train/test 比例,且需保证同类图片在 train/test 中分布均匀。项目用split_dataset.py实现分层抽样(stratified split),并生成train.txt/val.txt文件供自定义 DataLoader 使用(兼容旧版代码)。

# split_dataset.py import os, shutil, random, json from sklearn.model_selection import train_test_split def create_split_files(dataset_root: str, train_ratio: float = 0.8): classes = [d for d in os.listdir(dataset_root) if os.path.isdir(os.path.join(dataset_root, d))] train_list, val_list = [], [] for cls in classes: cls_path = os.path.join(dataset_root, cls) all_imgs = [os.path.join(cls, f) for f in os.listdir(cls_path) if f.lower().endswith(('.jpg', '.jpeg', '.png'))] # 分层抽样,确保每类比例一致 train_imgs, val_imgs = train_test_split( all_imgs, train_size=train_ratio, random_state=42, shuffle=True ) train_list.extend([(img, cls) for img in train_imgs]) val_list.extend([(img, cls) for img in val_imgs]) # 写入 train.txt: relative_path class_name with open("train.txt", "w", encoding="utf-8") as f: for img_rel, cls in train_list: f.write(f"{img_rel} {cls}\n") with open("val.txt", "w", encoding="utf-8") as f: for img_rel, cls in val_list: f.write(f"{img_rel} {cls}\n") # 生成类别 JSON 映射(供推理脚本读取) class_to_idx = {cls: idx for idx, cls in enumerate(classes)} with open("classes.json", "w", encoding="utf-8") as f: json.dump(class_to_idx, f, ensure_ascii=False, indent=2) print(f"完成划分:训练集 {len(train_list)} 张,验证集 {len(val_list)} 张") print(f"类别映射已保存至 classes.json") create_split_files("dataset", train_ratio=0.8)

注意classes.json是关键文件,内容形如{"钢铁侠": 0, "美国队长": 1, "雷神": 2}。训练脚本会自动读取该文件生成num_classes=3,无需手动修改模型配置。若新增类别,只需重新运行split_dataset.py即可。

3. Swin-Transformer 训练全流程:参数调整、中文日志与精度监控

本项目基于timm库(PyTorch Image Models)实现 Swin-Transformer,选用swin_tiny_patch4_window7_224(轻量级,适合单卡训练)。训练脚本train.py封装了完整的分布式训练逻辑,但默认以单卡模式运行,无需修改即可启动。

3.1 训练命令与核心参数说明

执行以下命令启动训练(假设已安装timm==0.9.2torch>=2.0):

python train.py \ --model swin_tiny_patch4_window7_224 \ --data-dir dataset \ --train-split train.txt \ --val-split val.txt \ --num-classes 3 \ --epochs 30 \ --batch-size 32 \ --lr 1e-4 \ --weight-decay 0.05 \ --opt adamw \ --sched cosine \ --warmup-epochs 5 \ --output results/swin_tiny_chinese \ --log-wandb \ --log-interval 20
参数说明推荐调整场景
--lr初始学习率小数据集(<500)建议5e-5~1e-4;若 loss 不降,尝试3e-5
--batch-size每卡 batch sizeRTX 3090 可设32;24G 显存卡建议16,避免 OOM
--warmup-epochswarmup 轮数防止 early stage 梯度爆炸,固定5即可
--sched cosine学习率调度器比 step decay 更稳定,收敛更快
--log-wandb启用 Weights & Biases 日志需提前pip install wandbwandb login

提示--data-dir dataset指向上一节生成的 clean 数据集根目录;--train-split--val-split指向split_dataset.py生成的train.txt/val.txt,格式为相对路径 类别名,完美支持中文类别名。

3.2 中文日志与训练过程可视化

train.py内置中文日志模块,关键信息自动转为中文输出(如训练轮次 15/30验证精度 0.912)。同时集成matplotlib绘图,每 epoch 结束后生成results/swin_tiny_chinese/train_log.png,包含:

  • 训练 loss(蓝色曲线)与验证 loss(橙色曲线)
  • 训练 acc(绿色)与验证 acc(红色)
  • 学习率变化(灰色虚线)
# 片段:绘图逻辑(train.py 内) def plot_training_log(logs: dict, save_path: str): fig, axes = plt.subplots(2, 1, figsize=(10, 8)) epochs = logs['epoch'] # Loss 曲线 axes[0].plot(epochs, logs['train_loss'], label='训练 Loss', color='blue') axes[0].plot(epochs, logs['val_loss'], label='验证 Loss', color='orange') axes[0].set_ylabel('Loss') axes[0].legend() axes[0].grid(True) # Accuracy 曲线 axes[1].plot(epochs, logs['train_acc'], label='训练 Acc', color='green') axes[1].plot(epochs, logs['val_acc'], label='验证 Acc', color='red') axes[1].set_xlabel('Epoch') axes[1].set_ylabel('Accuracy') axes[1].legend() axes[1].grid(True) plt.tight_layout() plt.savefig(save_path, dpi=300, bbox_inches='tight') plt.close()

注意:若plt中文显示为方块(常见于 Linux 服务器),需在绘图前插入:

import matplotlib matplotlib.rcParams['font.sans-serif'] = ['SimHei', 'DejaVu Sans', 'Arial Unicode MS'] matplotlib.rcParams['axes.unicode_minus'] = False

此设置已内置在train.py开头,确保图表中文正常渲染。

3.3 关键训练技巧:冻结 backbone 与混合精度

对于小数据集(如本项目的 347 张),直接微调全部参数易过拟合。项目提供--freeze-backbone选项,仅训练最后的 classifier head:

python train.py --freeze-backbone --lr 1e-3 ... # head 学习率可更高

同时启用--amp(Automatic Mixed Precision)加速训练并节省显存:

# 训练循环中(train.py) scaler = torch.cuda.amp.GradScaler() # 初始化 scaler for data, target in train_loader: optimizer.zero_grad() with torch.cuda.amp.autocast(): # 自动混合精度前向 output = model(data) loss = criterion(output, target) scaler.scale(loss).backward() # 缩放梯度 scaler.step(optimizer) # 更新参数 scaler.update() # 更新 scaler

实测开启--amp后,RTX 3090 单卡训练速度提升 1.8 倍,显存占用降低 35%。

4. 推理与部署:中文类别输出、批量预测与错误分析

训练完成后,模型权重保存在results/swin_tiny_chinese/checkpoint.pth。推理脚本inference.py支持两种模式:单图预测(调试用)和批量预测(生产用),所有输出均保留中文类别名。

4.1 批量预测:自动处理 inference/ 下所有图片

将待预测图片放入inference/目录(支持子目录),运行:

python inference.py \ --model-path results/swin_tiny_chinese/checkpoint.pth \ --classes-json classes.json \ --input-dir inference/ \ --output-dir results/predictions/ \ --top-k 2

脚本会递归扫描inference/下所有.jpg/.jpeg/.png文件,对每张图输出 top-2 预测结果(含概率),生成results/predictions/predictions.csv

filename,rank_1_class,rank_1_prob,rank_2_class,rank_2_prob inference/ironman_001.jpg,钢铁侠,0.923,雷神,0.041 inference/captain_002.jpg,美国队长,0.876,钢铁侠,0.089
# inference.py 核心逻辑 def predict_batch(model, transform, classes_json, input_dir, output_csv): with open(classes_json, 'r', encoding='utf-8') as f: class_idx = json.load(f) idx_to_class = {v: k for k, v in class_idx.items()} # 反向映射 results = [] for img_path in Path(input_dir).rglob("*.*"): if img_path.suffix.lower() not in ['.jpg', '.jpeg', '.png']: continue try: img = Image.open(img_path).convert('RGB') img_tensor = transform(img).unsqueeze(0).to(device) # 加 batch 维度 with torch.no_grad(): logits = model(img_tensor) probs = torch.nn.functional.softmax(logits, dim=1) top_probs, top_indices = torch.topk(probs, k=2) row = { 'filename': str(img_path.relative_to(input_dir)), 'rank_1_class': idx_to_class[top_indices[0, 0].item()], 'rank_1_prob': f"{top_probs[0, 0].item():.3f}", 'rank_2_class': idx_to_class[top_indices[0, 1].item()], 'rank_2_prob': f"{top_probs[0, 1].item():.3f}" } results.append(row) except Exception as e: print(f"[FAIL] {img_path} 预测失败: {e}") pd.DataFrame(results).to_csv(output_csv, index=False, encoding='utf-8-sig')

注意encoding='utf-8-sig'确保 Excel 能正确打开 CSV 中的中文。idx_to_classclasses.json动态构建,新增类别无需改代码。

4.2 错误分析:定位低置信度预测与类别混淆

高精度(0.91)不等于无问题。项目提供analyze_errors.py,自动统计:

  • 低置信度样本:top-1 概率 < 0.7 的图片,存入error_analysis/low_confidence/
  • 混淆矩阵:生成confusion_matrix.png,可视化各类别间误判情况;
  • 典型错误案例:提取每类被误判为其他类的 top-3 图片,便于人工复核。
# analyze_errors.py 片段 def generate_confusion_matrix(y_true, y_pred, class_names, save_path): cm = confusion_matrix(y_true, y_pred) plt.figure(figsize=(8, 6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.title('混淆矩阵') plt.ylabel('真实类别') plt.xlabel('预测类别') plt.tight_layout() plt.savefig(save_path, dpi=300, bbox_inches='tight') plt.close()

运行后,error_analysis/confusion_matrix.png显示:

  • 钢铁侠 → 美国队长 误判 3 次(因盔甲色调相似)
  • 雷神 → 钢铁侠 误判 5 次(因闪电特效与金色装甲混淆)
    此分析直接指导数据增强策略:对钢铁侠/美国队长添加更多光照变化,对雷神增加闪电背景合成。

5. 进阶技巧:中文标签可视化与模型轻量化部署

当模型进入业务系统,需解决两个实际问题:一是预测结果需嵌入中文 UI(如 Web 页面 tooltip),二是边缘设备(如 Jetson Orin)需压缩模型体积。本节提供即插即用方案。

5.1 中文标签热力图:Grad-CAM 可视化聚焦区域

gradcam_visualize.py基于captum库,为任意输入图生成中文类别对应的热力图,直观显示模型关注区域:

python gradcam_visualize.py \ --model-path results/swin_tiny_chinese/checkpoint.pth \ --classes-json classes.json \ --input-image inference/ironman_001.jpg \ --output-dir results/gradcam/ \ --target-class 钢铁侠

输出results/gradcam/ironman_001_钢铁侠_cam.jpg,叠加在原图上,红色区域即模型判定“钢铁侠”的依据(如面部、胸口反应堆)。代码自动匹配target-classclasses.json中的索引,无需手动查数字。

# gradcam_visualize.py 核心 from captum.attr import LayerGradCam from captum.attr import visualization as viz def visualize_gradcam(model, input_tensor, target_class, class_to_idx, save_path): model.eval() target_idx = class_to_idx[target_class] # 自动查中文名对应索引 layer_gc = LayerGradCam(model, model.layers[-1].blocks[-1].norm1) # Swin 最后一个 block 的 norm cam = layer_gc.attribute(input_tensor, target=target_idx) cam = cam[0].cpu().detach().numpy() # (C, H, W) -> (H, W) # 可视化 original_img = np.array(Image.open(input_tensor_path).convert('RGB')) viz.visualize_image_attr( cam, original_img, method='heat_map', sign='positive', show_colorbar=True, title=f"Grad-CAM for {target_class}", plt_fig_axis=None, use_pyplot=True ) plt.savefig(save_path, dpi=300, bbox_inches='tight') plt.close()

5.2 模型导出为 TorchScript:兼容无 Python 环境

生产环境常需脱离 Python 解释器。export_model.py将训练好的 Swin 模型导出为.pt格式,可在 C++/Java 中加载:

python export_model.py \ --model-path results/swin_tiny_chinese/checkpoint.pth \ --classes-json classes.json \ --output-path models/swin_tiny_chinese.pt

导出的模型已封装预处理(resize、normalize)和后处理(softmax、top-k),调用示例(C++):

#include <torch/script.h> auto module = torch::jit::load("models/swin_tiny_chinese.pt"); std::vector<torch::jit::IValue> inputs; inputs.push_back(image_tensor); // shape: [1,3,224,224], dtype: float32 at::Tensor output = module.forward(inputs).toTensor(); // output[0] 为类别概率向量,索引对应 classes.json 顺序

提示:导出前自动插入torch.jit.script装饰器,确保所有控制流(如if)可追踪。classes.json被打包进模型,output[0]的第 0 位永远对应classes.json中第一个中文类别。

5.3 量化压缩:INT8 模型体积减少 75%,精度仅降 0.01

swin_tiny进行动态量化(Dynamic Quantization),无需校准数据集:

# export_model.py 中量化逻辑 quantized_model = torch.quantization.quantize_dynamic( model, {torch.nn.Linear, torch.nn.Conv2d}, dtype=torch.qint8 ) torch.jit.save(torch.jit.script(quantized_model), "models/swin_tiny_quantized.pt")

量化后模型体积从 127MB → 31MB,CPU 推理速度提升 2.3 倍。在本项目数据集上,验证精度从 0.910 → 0.909,可忽略不计。此步骤已集成进export_model.py,添加--quantize参数即可启用。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/15 3:16:10

Django影评社区开发实战:从数据建模到生产部署全记录

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/15 3:15:02

SpringBoot与Maven构建智慧社区报修平台实战

1. 项目概述&#xff1a;智慧社区报修平台的技术选型智慧社区作为现代城市管理的重要单元&#xff0c;其报修平台的搭建需要兼顾快速开发与稳定运行的双重需求。SpringBoot作为当前Java领域最主流的微服务框架&#xff0c;其"约定优于配置"的理念能显著降低开发门槛。…

作者头像 李华
网站建设 2026/9/15 3:14:52

YOLOv5道路破损检测实战:数据集准备、模型训练与推理优化

简介&#xff1a;面向道路破损检测场景的YOLOv5完整资源包&#xff0c;兼顾算法学习与工程项目落地。包内集成训练好的YOLOv5权重&#xff0c;可直接用于图片或视频中的道路破损推理&#xff1b;配套7000余张真实场景道路破损图片&#xff0c;已用LabelImg标注为VOC与YOLO两种格…

作者头像 李华
网站建设 2026/9/15 3:14:50

飞机目标检测数据集:VOC+YOLO双格式小目标优化实践

简介&#xff1a;本资源是一份专为计算机视觉目标检测任务设计的高质量飞机图像数据集&#xff0c;适用于深度学习初学者、算法工程师及科研人员开展模型训练与验证。数据集共7931张JPG图像&#xff0c;全部配有Pascal VOC格式XML标注文件与YOLO格式TXT标注文件&#xff0c;类别…

作者头像 李华
网站建设 2026/9/15 3:14:33

车规级ECU基于UDS协议的CAN OTA升级实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华