资讯动态

基于ResNet的遥感图像分类实战:PyTorch实现有湖无湖二分类

发布时间:2026/9/28 22:41:26 来源:尧图企业网站定制
简介这套基于PyTorch的ResNet图像分类工程专为遥感领域初学者设计用于识别遥感影像中有湖泊与无湖泊两类场景属于典型的二分类应用。整个包仅191KB共7个文件3个Python脚本分别负责生成数据列表、训练CNN模型、调用PyQt界面完成预测展示2张jpg示例图提示图片放置位置1个requirements.txt列出依赖库1份docx说明文档则详细梳理从环境搭建到跑通训练的每一步。代码最大特点是每行都有中文注释即使是零基础用户也能按注释逐步理解模型构建、数据加载和训练流程由于不含数据集图片下载后需自行收集图片按类别放入文件夹若有更多分类需求也可新建文件夹扩展为多分类任务灵活性较高。推荐使用Anaconda安装Python 3.7/3.8与PyTorch 1.7.1/1.8.1依赖列表已备好目前已有150人学习下载适合想快速上手遥感图像分类的PyTorch新手作为起点。1. 遥感图像分类实战起点ResNet 识别有无湖泊这套带逐行注释的代码能直接改着用ResNet 图像分类在遥感场景里是出现频率最高的起步方案。这个压缩包把「遥感影像里有没有湖泊」这个二分类任务拆成三个可独立运行的 py 文件一个生成训练列表、一个训练 ResNet、一个用 PyQt 界面做推理而且每一行都带中文注释。我拿到手第一感觉是「终于有个不用靠猜来理解的落地项目了」它不依赖现成数据集分类文件夹可以自己造连类别数量都能随手改。适合刚装好 PyTorch 想跑通第一个完整图像分类链路的人也适合想快速验证深度学习到底能不能把遥感图分明白的从业者。下面按实际跑通顺序拆开讲。2. 先理顺三件套的配合数据流顺序、目录结构与环境对齐2.1 三个 py 文件的分工谁生成输入、谁训练、谁做界面打开压缩包第一眼是三个 py 文件加一份说明文档。很多人第一次拿到这种项目会先双击 02CNN训练数据集.py 看看能不能跑这其实是常见的第一步弯路训练脚本需要读取 01 生成的 txttxt 还没生成脚本在数据加载阶段就会报错。正确顺序是先读说明文档搭好数据集目录然后从 01 开始逐步推进。三个脚本的分工可以先用一张表说清楚文件职责运行时机01生成txt.py扫描数据集目录生成「图片路径 类别编号」的 train.txt每次数据集变动后第一个运行02CNN训练数据集.py读取 train.txt加载图片并训练 ResNet保存模型权重等 01 生成的 txt 确认无误后运行03pyqt界面.py启动 PyQt 窗口选图调用模型输出类别与置信度训练完成拿到权重之后运行这张表最关键的是第二列和第三列的组合01 生成的东西是 02 的输入02 训练出的权重是 03 的输入链路单向。我见过有人把 03pyqt界面.py 单独拎出来研究很久最后发现报错原因是模型根本没训练过——界面代码没有权重文件可加载这属于典型的「卡在中间不知道前后依赖」的问题。为什么要拆成三个文件而不是一个脚本跑完好处是每一步能单独验证。01 跑完打开 train.txt 检查路径和标签是否对应02 跑完看 loss 曲线判断是继续训练还是该停03 只负责推理界面出问题时不需要怀疑是训练流程的锅。这种「一个任务一个文件」的拆分习惯对刚入门的人来说比一个大而全的脚本友好得多报错位置清晰排查范围小每一行还能顺着逻辑读下去。2.2 数据集目录约定没有自带图片分类文件夹可以自己搭压缩包刻意不含数据集图片这让一部分人一开始有点慌。实际上这恰恰是这套代码的灵活之处类别目录由你自己创建「是否有湖泊」只是举例你想做「有无云层」「有无植被」都可以只需要换掉文件夹名字和里面的图片。默认的目录结构是这样数据集/ ├── 无湖泊/ │ ├── 1.jpg │ ├── sample_01.jpg │ └── sample_02.jpg └── 有湖泊/ ├── 1.jpg └── lake_01.jpg每个类别文件夹里那张 1.jpg 是提示图告诉你这类图片该放哪儿。它本身是一张真图片扩展名也是 jpg所以 01生成txt.py 扫描文件夹时不会区分「提示图」和「训练图」直接就把 1.jpg 也写进 train.txt——等于训练集里混入了一张你没打算喂给模型的样本。我每次整理数据时都会先把提示图移出数据集目录再重新生成 txt避免这种隐性干扰。类别文件夹的名字不是固定的增加分类也很直接新建一个文件夹把图片放进去重跑 01训练脚本的类别数会跟着目录数自动变化。但有一个前提必须说清楚类别越细需要的样本量越大。二分类每个类 50 张图勉强能跑三分类每类最好 100 张以上否则少数类几乎学不出来。遥感图本身类内差异很大——同样是「有湖泊」太湖和一个小水塘在形态上差很多——样本太少模型很容易只记住其中一种形态。图片收集时我会留意一个原则训练样本要覆盖不同季节、不同天气、不同分辨率的遥感影像。湖泊在影像上一般是深色块但阴影、深色农田、云层阴影都会造成类似信号样本多样性不够模型就会把这些干扰项一起学进去。这也是「无湖泊」类里最难构造的部分你需要的不是随便凑几十张图而是尽量多的「看着像湖但实际不是」的困难样本。数据来源可以用公开遥感数据源或项目自采航拍图注意版权和分辨率一致性图像标注这一步直接决定模型上限。2.3 环境对齐Anaconda 建环境Python 3.8 与 PyTorch 1.8.1 的搭配理由说明文档里推荐 Anaconda 安装 Python 3.7 或 3.8PyTorch 用 1.7.1 或 1.8.1。这个组合放在今天不是最新但在 Windows 上属于非常成熟稳定的搭配这两个 Python 版本对 torchvision 自带的 ResNet 预训练权重兼容性很好不会遇到算子缺失PyTorch 1.7/1.8 的 CUDA 支持稳定网上教程和踩坑案例都多遇到问题容易搜到答案而且这套代码本身也没有用到 torch 2.x 才有的新特性追新没有意义。我的环境搭建步骤是conda create -n resnet-lake python3.8 conda activate resnet-lake pip install -r requirement.txtrequirement.txt 是依赖清单包含 torch、torchvision、opencv-python、PyQt5 等常见库。注意一点torch 的安装包很大pip 直接装默认装 CPU 版本如果你有 N 卡建议去 PyTorch 官网按 CUDA 版本生成安装命令再装训练速度会有数量级差别。如果设备没有独显CPU 版本也能跑这套流程只是训练 20 个 epoch 的时间会从几分钟变成几十分钟属于可以接受但需要耐心的范围。环境安装最大的坑是包管理器混用。conda 和 pip 对依赖的解析逻辑不同同一个包先后用两个源装很可能会出现「import torch 成功import torchvision 版本报错」的怪问题。我的习惯是conda 只负责创建环境装依赖一律用 pip全程不回头。这样至少能保证依赖关系是同一套体系解析的排错时少一个变量。3. 跑通训练闭环01生成txt.py 与 02CNN训练数据集.py 逐段实操3.1 运行 01生成txt.py把分类文件夹转成训练清单01生成txt.py 做的事一句话能说清遍历数据集目录把每张图片的路径和它所属类别的编号写进 train.txt。这是训练流程的地基因为 PyTorch 的 Dataset 通常从文本文件读样本列表而不是直接读文件夹。代码的核心逻辑长这样import os dataset_path 数据集 # 改成你自己的实际路径 txt_path train.txt with open(txt_path, w, encodingutf-8) as f: for class_idx, class_name in enumerate(os.listdir(dataset_path)): class_dir os.path.join(dataset_path, class_name) if not os.path.isdir(class_dir): continue for img_name in os.listdir(class_dir): img_path os.path.join(class_dir, img_name) f.write(f{img_path} {class_idx}\n)写得比较直白enumerate 在遍历文件夹列表时同时生成类别编号第一个文件夹编号 0第二个编号 1依次类推往里再扫一层把每张图片的相对路径和编号拼成一行写入 txt。这里没做扩展名过滤实际工程里我一般会在写入前加一个判断只保留 jpg、png、bmp 这些常见图片格式省得把文件夹里不小心混入的临时文件也写进去。生成出来的 train.txt 内容是这种格式数据集/无湖泊/sample_01.jpg 0 数据集/有湖泊/lake_01.jpg 1行尾的数字就是类别编号。02 训练脚本读取时根本不看文件夹名只认这个数字。所以编号规则看起来随意实际上有一个隐患如果 01 在某个版本里遍历顺序变化或者你中途新增了一个文件夹编号就可能和以前不对应。我的习惯是数据集一旦改动就重新跑一遍 01绝不手工去改 txt——手工改错位的代价是训练出来的模型标注整体错乱排查起来非常痛苦。3.2 训练脚本主流程自定义 Dataset、ResNet 模型与训练循环02CNN训练数据集.py 是三个文件里的核心它做的事情可以拆成四步读 txt 构建数据集、定义 ResNet 模型、写训练循环、保存权重。第一步是数据入口import torch from torch.utils.data import Dataset from PIL import Image class ImageListDataset(Dataset): def __init__(self, txt_path, transformNone): self.samples [] with open(txt_path, r, encodingutf-8) as f: for line in f: img_path, label line.strip().split() self.samples.append((img_path, int(label))) self.transform transform def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label self.samples[idx] img Image.open(img_path).convert(RGB) if self.transform: img self.transform(img) return img, label这段 Dataset 负责把 train.txt 的每一行还原成「图片、标签」对。getitem里有个人容易忽略的细节convert(RGB)。遥感截图常以 PNG 格式保存而 PNG 可能是 RGBA 四通道ResNet 的第一层卷积只接受三通道输入。不做这一步转换训练会在数据加载阶段直接报 expected 3 input channels, got 4。所以这句话宁可留着哪怕你确定自己的图片都是 JPG它也不会产生副作用。注意convert(RGB) 必须放在getitem里而不是只在数据准备阶段做一次。训练中途才炸出来的通道错误绝大多数都是因为偷懒只转了部分图片。模型创建和训练循环是重点import torchvision.models as models import torch.nn as nn import torch.optim as optim device torch.device(cuda if torch.cuda.is_available() else cpu) model models.resnet18(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) model.fc nn.Linear(model.fc.in_features, 2) # 二分类改成你的类别数 model model.to(device) optimizer optim.Adam(model.parameters(), lr1e-4) criterion nn.CrossEntropyLoss() for epoch in range(epochs): model.train() running_loss 0.0 for images, labels in train_loader: images, labels images.to(device), labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() # 每个 epoch 存一次训练中断也不至于白跑 torch.save(model.state_dict(), best_model.pth) print(fEpoch {epoch1}/{epochs}, Loss: {running_loss/len(train_loader):.4f})weightsResNet18_Weights.IMAGENET1K_V1 会自动下载 ImageNet 预训练权重。这是遥感小数据集能训练出结果的关键随机初始化的 ResNet 需要海量数据才能学到有效特征而预训练权重相当于给模型一个「已经会识别纹理和边缘」的起点你只需要在它基础上微调。model.fc 那行把原来 1000 类的全连接层换成 2 类输出训练时真正需要大改的也就是这一层。训练循环里几个操作要说一下optimizer.zero_grad() 清空上一轮累积的梯度loss.backward() 计算当前 batch 的梯度optimizer.step() 更新权重。漏掉 zero_grad 是最常见的低级错误——梯度会跨 batch 累积loss 曲线看起来在降实际权重更新方向是乱的。每个 epoch 保存一次模型是防中断的土办法训练到 48 轮崩了至少还有第 47 轮的权重不用从头再来。3.3 关键训练参数起步值怎么设什么时候该动它我在这套场景下的起步参数如下参数起步值调整方向batch_size8 或 16显存不足先降 batch不建议降图片尺寸learning_rate1e-4训练不稳就降到 1e-5很少需要往上调epochs20~50观察验证 loss连续 5 轮不降就提前停图片输入尺寸224x224ResNet 的标准输入torchvision 内部会做缩放optimizerAdam小数据集上比 SGD 好收敛不用调动量参数batch_size 和显存是直接捆绑的。遥感影像普遍分辨率高内存里加载几十张没问题但一进 GPU 显存就吃紧。我的经验是先按 16 跑爆显存就降到 8还爆就查是不是图片太大没做 Resize。学习率的直觉很重要预训练模型已经处于一个相对低的 loss 区域所以学习率应该比从零训练小一个数量级。1e-3 在这种场景下偏大loss 可能震荡1e-4 稳妥1e-5 保守。如果 loss 降得很慢先别怀疑学习率去检查数据对不对。4. 训练避坑与常见问题排查四个高频翻车现场及处理记录4.1 准确率卡在 50%二分类模型的玄学临界点现象训练了十几个 epoch训练集准确率一直在 50% 上下波动和随机猜测几乎一样。二分类的 50% 是一个特别容易让人产生自我怀疑的点模型到底没学到东西还是数据出问题了原因排查顺序有三个——学习率太大导致 loss 震荡类别编号错位txt 里的标签和实际文件夹顺序不对应两类样本数量严重不均衡少数类被模型直接忽略。学习率太大的情况下预训练权重的 loss 会因为一次大步长更新被冲高然后又慢慢降回来整体看起来就像在 50% 附近抖动。解决先把学习率降到 1e-5 重训一个短实验排除优化器问题再用下面的脚本检查标签分布from collections import Counter labels [] with open(train.txt, r, encodingutf-8) as f: for line in f: labels.append(int(line.strip().split()[1])) print(Counter(labels))如果 Counter 输出显示某一类只有个位数样本问题不在模型在数据。二分类任务里我不建议用加权损失来硬调少数类先加样本更实际。如果标签分布均衡、学习率也正常还是卡 50%就去查 01 生成的 txt 里「图片路径」是否真的指向正确的类——错位的标注会让模型学到完全相反的东西准确率甚至会低于 50%。4.2 报错 expected 3 input channels, got 4遥感切片常带透明通道现象训练跑进 DataLoader 加载阶段报通道数不一致错误指向某一张 PNG 图片。最折磨人的是它有时候跑到一半才炸当时还以为是自己数据路径写错了。原因遥感截图经常是 RGBA 四通道 PNG多出来的 A 通道是透明度信息。ResNet 第一层卷积核是三通道输入四通道直接报错。有的图集中只有个别几张是四通道但 DataLoader 逐张加载跑了几百轮才炸一次。A 通道在很多软件里是恒定值对模型没有信息量但它的存在会直接改变输入张量的形状。解决在 Dataset 的getitem里统一 convert(RGB)。更稳妥的做法是在收集数据阶段就跑一遍批量转换把数据集里所有 PNG 统一转成 RGB JPG训练前就消除隐患。我建议两种都做转换脚本用于数据整理convert(RGB) 作为代码兜底。转换脚本也很简单用 PIL 打开后 save 成 jpg 就行关键是要递归遍历所有子目录。4.3 加载权重报错DataParallel 的 module. 前缀与 torchvision 版本差异现象加载 02 保存的权重时报 Missing key(s) in state_dict 或 Unexpected key(s)训练和加载的代码看起来一模一样但就是加载不进去。原因如果权重是用 DataParallel 包装过的模型保存的state_dict 里所有键名都带 module. 前缀torchvision 不同小版本之间ResNet 的层命名也可能有细微差异。这两类问题都会导致 key 对不上。PyTorch 版本从 1.x 升到 2.x 时torchvision 对 ResNet 的实现细节有过调整同样一个 resnet18两个版本导出的权重文件并不保证完全通用。解决写一个通用加载函数去掉前缀再加载state_dict torch.load(best_model.pth, map_locationcpu) new_state_dict {} for k, v in state_dict.items(): if k.startswith(module.): k k[7:] # 去掉 module. 前缀 new_state_dict[k] v model.load_state_dict(new_state_dict, strictFalse)strictFalse 不是让你闭着眼忽略不匹配而是先加载看看哪些 key 没对上再决定是修前缀还是修模型结构。这个函数我固化在代码模板里所有项目的权重加载都走它省掉了大量排查时间。如果你发现去掉前缀还是报错那就把打印出来的 key 列表和当前模型的 key 列表摆在一起 diff一眼就能看出是哪一层的命名差异。4.4 训练指标好看单张预测翻车训练和推理的预处理不一致现象验证集准确率 90% 以上但用 03pyqt界面.py 实测一张没见过的遥感图结果错了而且是稳定地错。换一张图又是稳定地错这种感觉特别像模型训练失败了但看指标又没问题。原因训练时用的 transform 和推理时用的 transform 不一致最常见的是归一化的均值和标准差不同或者训练时有随机翻转、推理时没有对齐。ResNet 对输入分布很敏感归一化参数差一点预测概率就会整体偏移。另一个隐蔽点是 Resize 和目标尺寸不一致训练时按短边缩放、推理时直接拉伸图像的几何形变完全不同。解决把训练时的完整 transform 原样复制到预测脚本包括 Resize、Normalize 的均值和标准差。03pyqt界面.py 里的预处理必须和 02 训练时保持一致这是推理代码最容易出错的点也是排查优先级最高的点。我那一次翻车就是改代码时顺手把 Normalize 的均值写错了概率输出全部偏向一类当时还以为是模型不行冷静下来对比两个 transform 才发现问题。要验证预处理是否真的对齐最直接的方法是拿训练集里的同一张图分别走训练代码和推理代码打印预处理后的像素统计值两个结果必须完全一致。5. 让模型出结果03pyqt界面.py 的推理实现与多分类改造5.1 从模型权重到界面按钮PyQt 推理链路拆解03pyqt界面.py 把上一步训好的模型封装成一个可视化窗口。PyQt 界面的核心职责就两个选图、显示结果。模型加载和预处理都在后台完成。一个简化的核心结构如下import sys import torch from PyQt5.QtWidgets import (QApplication, QWidget, QPushButton, QLabel, QFileDialog) class LakeWindow(QWidget): def __init__(self): super().__init__() self.model self.load_model() # 启动时加载一次 self.btn QPushButton(选择图片, self) self.btn.clicked.connect(self.predict_image) self.result_label QLabel(请选择一张遥感影像, self) def load_model(self): # 与训练阶段相同的模型结构二分类 model models.resnet18() model.fc nn.Linear(model.fc.in_features, 2) model.load_state_dict(torch.load(best_model.pth, map_locationcpu)) model.eval() return model def predict_image(self): path, _ QFileDialog.getOpenFileName( self, 选择图片, , Images (*.jpg *.png)) if not path: return prob self.run_inference(path) # 返回长度为2的概率列表 label_idx prob.index(max(prob)) labels [无湖泊, 有湖泊] self.result_label.setText(f{labels[label_idx]}置信度 {prob[label_idx]:.2f})界面代码有两个值得注意的工程习惯。一是模型只在窗口初始化时加载一次不要在每次点按钮时重新 load_state_dict——那样每点一次图就要等几秒体验极差。二是 load_state_dict 之后必须调用 model.eval()否则模型里的 BatchNorm 层会用训练模式计算同一张图每次推理结果都可能不一样。eval() 这行漏掉的后果很隐蔽偶尔错一次很难复现容易被误判成模型问题。推理函数本身的逻辑和训练时前向传播一致图片读入 → 预处理 → 升维 → no_grad 前向 → softmax。升维用 unsqueeze(0) 是因为模型要求输入带 batch 维单张图本来就是 3 维张量通道、高、宽加一个 0 维度变成 1,3,224,224 才能进模型。softmax 把最终输出转成概率分布方便界面显示置信度。5.2 从二分类到多分类改三个地方连带检查两件事这套代码的分类数量不是写死的。把「有/无湖泊」扩成「有湖泊 / 无湖泊 / 含云层」三分类只需要改三个位置数据集文件夹加一个「含云层」目录并放入图片重跑 01生成txt.py 让 txt 里的类别编号更新把 02 和 03 里输出维度从 2 改成 3界面标签数组同步改# 模型输出层 model.fc nn.Linear(model.fc.in_features, 3) # 界面标签数组顺序必须和数据集目录顺序一致 labels [无湖泊, 有湖泊, 含云层]改完之后有两个连带检查。第一txt 里的编号顺序必须和界面标签数组顺序一致模型输出第 0 维对应「无湖泊」界面标签数组第 0 个就必须是「无湖泊」错位的话模型准了界面也显示错。第二新增类别的样本量要撑得起训练三分类里「含云层」如果只有 5 张图模型大概率直接忽略这个类softmax 输出里它的概率永远接近 0。遥感影像分类里类别不平衡是常态我通常先给少数类补图实在补不到再考虑用类别权重。多分类改造还有一个容易忽略的点训练脚本里如果用的是 DataLoader 的 shuffleTrue每个 epoch 样本顺序都会乱这是正常的。但验证时不要 shuffleshuffle 只用于训练阶段打乱顺序防止模型记住样本出现顺序。03 的推理是一次一张图不涉及这个问题但如果你把 02 的 DataLoader 代码复制去写验证脚本记得把 shuffle 参数关掉。5.3 界面扩展思路批量预测与结果导出如果不想一张张点常见做法是给 03 增加一个「批量预测」按钮逻辑是遍历一个文件夹下的所有图片循环调用 run_inference把结果写入 CSV。这个改造不复杂但价值很大——尤其你现在需要自己搜集数据、自己评估模型批处理能把 50 张测试图的结果一次跑出来。import csv with open(predict_results.csv, w, newline, encodingutf-8) as f: writer csv.writer(f) writer.writerow([图片, 类别, 置信度]) for img_file in test_images: prob run_inference(img_file) label_idx max(range(len(prob)), keylambda i: prob[i]) writer.writerow([img_file, labels[label_idx], f{prob[label_idx]:.2f}])批量预测配合结果 CSV能让你快速统计一整批图的准确率而不是靠肉眼一张张看界面。做遥感分类评估时我一般都会先跑一遍批量模式把分布情况摸清楚再回头决定要不要调整训练策略。6. 验证模型的土办法准确率之外先看类别概率分布模型训完一般人的第一反应是看验证集准确率。但在「没有标准测试集、需要自己搜集图片」的条件下准确率很容易骗人——你搜集图片时的主观筛选会让验证集和训练集长得太像模型其实记住了你搜集图片的风格而不是湖泊特征。我自己的习惯是看概率分布。做法很简单挑 20 张有湖泊和 20 张无湖泊的图逐张记录 softmax 输出。把每张图的预测概率打印出来with torch.no_grad(): outputs model(input_tensor) prob torch.softmax(outputs, dim1).squeeze().tolist() print(f无湖泊 {prob[0]:.2f} | 有湖泊 {prob[1]:.2f})区分信号就一条健康的模型多数预测概率应该落在 0.7 以上正负样本的分布明显分开如果大量预测集中在 0.50.6说明模型学到的是模糊的纹理统计没有形成清晰的特征换一批图很可能翻车。概率分布是比准确率早一步暴露问题的手段——准确率 80% 但概率全是 0.55和准确率 80% 但概率大多是 0.9 的模型鲁棒性完全不是一个级别。接下来把分错的图单独挑出来看。湖泊在遥感图上通常是深色、边界平滑、形状不规则的水体如果模型把所有深色区域都判成有湖泊那它学到的特征其实是「暗色块 湖泊」而不是「水体 湖泊」。这种时候要去「无湖泊」类里补充困难样本深色的山体阴影、深色的农田、云层阴影它们都是「看着像湖但不是湖」的典型反例。ResNet 这类模型提取的主要是全局纹理特征对细粒度目标不敏感所以反例的质量直接决定模型能不能学会区分真正的湖泊。顺带说一句如果你后续想做的不是整图分类而是逐像素分割那就要换 SegFormer 那类语义分割模型那是另一套工程了。从那以后我每次准备交付模型前都会强制走一遍概率分布抽查挑图、跑批量预测、看分布、挑错误样本、补反例。整个过程花不了十分钟但能避免交付一个验证集好看、实测掉链子的模型。如果这套代码你也想拿去跑通建议从最小闭环开始先每个类 30 张图把流程完整走一遍确认链路没问题了再逐步加数据。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑