简介:本资源是一份面向高校计算机视觉方向本科生与毕业设计学生的遥感图像语义分割实践方案,聚焦U-Net网络在多光谱遥感影像中的建筑目标自动提取任务。资源完整复现了基于Inria航拍数据集的训练、验证与预测全流程,创新性地对比了标准交叉熵与自定义类别平衡交叉熵损失函数对F1 Score的影响(提升8.5%),为遥感领域小样本、类别不均衡场景提供可复用的改进思路。压缩包共69个文件,含6个核心Python训练/预测脚本(train.py、predict.ipynb等)、32张标注与预测结果PNG图、5个LaTeX论文章节(含chap1–chap5.tex及Bib参考文献)、1份完整毕业论文PDF及配套工具脚本(如start_jupyter.ps1、create_dataset.ipynb),总大小46.98MB,结构清晰,开箱即用。目前已有453人学习下载,适合需快速搭建遥感分割基线、理解U-Net适配遥感数据的关键调整点、并直接复用于课程设计或毕设答辩的中高级学习者。
1. 为什么毕业设计选 U-Net 做遥感图像语义分割,不是跟风,是真能跑通、能调参、能交稿
你手头有一份标注好的遥感影像——比如某市域的 GF-2 卫星图,带 5 类地物标签(水体、建筑、道路、林地、裸土),但用 OpenCV 阈值+形态学硬抠边界?结果边缘锯齿、小目标漏检、阴影区全崩。这时候翻毕业设计选题库,“U-Net 语义分割”高频出现,不是因为论文灌水多,而是它在小样本、高分辨率、弱纹理遥感图上,比 DeepLabv3+ 更稳、比 Mask R-CNN 更轻、比 FCN 更易收敛——尤其对本科生:不需要配 GPU 集群,一块 RTX 3060 就能从数据准备到模型部署走完闭环。我带过 17 届到 24 届共 32 个遥感方向毕设,用 U-Net 的 28 个按时答辩,失败的 4 个全是卡在数据格式没对齐、标签图通道错位、验证指标算错这三处。本文不讲论文怎么写,只拆解:怎么把.zip里那堆.tif和.png变成可训练的 PyTorch Dataset;U-Net 结构里哪 3 行代码决定你能否避开梯度爆炸;验证时为什么 IoU 看着高,但导出的分割图全是“马赛克块”;以及最关键的——如何用 1 个predict.py脚本,把训练好的模型直接拖进导师电脑双击运行。全文所有命令、参数、路径、报错截图,都来自真实毕设环境(Windows 10 + Python 3.9 + PyTorch 1.13)。
2. 从 .zip 解压到 PyTorch Dataset:遥感图像语义分割的数据准备四步法
遥感图像语义分割的数据准备,和自然图像有本质区别:卫星图是多光谱(常为 4 波段:R/G/B/NIR),标签图是单通道灰度(0~4 对应五类),且尺寸动辄 5000×5000 像素。直接喂进 U-Net 会 OOM,必须切块。但切块不是简单用cv2.resize,得保纹理、对齐坐标、防跨类割裂。下面四步是我在 2023 年指导 11 个毕设项目验证过的最小可行流程。
2.1 解压与目录结构标准化:拒绝“文件全塞根目录”玄学
拿到【毕业设计】基于 U-Net 网络的遥感图像语义分割 .zip,先别急着解压。打开压缩包看内部结构——常见三种坑:
- ❌
img/,label/,train.txt混在根目录(无 train/val/test 分离) - ❌ 标签图是 RGB 彩色 PNG(每个像素是 (255,0,0) 这种三元组,非单通道灰度)
- ❌ 图像名不一致(
IMG_001.tifvslabel_001.png)
正确做法:强制统一为以下结构
data/ ├── images/ # 所有原始遥感图,.tif 或 .png,命名 IMG_001.tif, IMG_002.tif... ├── labels/ # 所有标签图,单通道 uint8 灰度图,命名 IMG_001.png, IMG_002.png... └── splits/ # 划分文件,train.txt / val.txt / test.txt,每行一个图像名(不含扩展名)提示:如果原始 zip 里标签是彩色 PNG,用以下脚本批量转单通道(关键:用
cv2.IMREAD_GRAYSCALE读,再np.where映射颜色到类别 ID):
# convert_labels.py import cv2 import numpy as np import os # 定义颜色到类别ID的映射(按你的数据集实际颜色填) color_to_id = { (0, 0, 0): 0, # 黑色→背景 (255, 0, 0): 1, # 红色→建筑 (0, 255, 0): 2, # 绿色→林地 (0, 0, 255): 3, # 蓝色→水体 (255, 255, 0): 4 # 黄色→道路 } def rgb_to_gray_label(rgb_path, gray_path): rgb = cv2.imread(rgb_path) h, w, _ = rgb.shape gray = np.zeros((h, w), dtype=np.uint8) for i in range(h): for j in range(w): pixel = tuple(rgb[i, j]) gray[i, j] = color_to_id.get(pixel, 0) # 未定义颜色默认为0 cv2.imwrite(gray_path, gray) # 批量转换 for rgb_file in os.listdir("labels_rgb"): if rgb_file.endswith(".png"): rgb_path = os.path.join("labels_rgb", rgb_file) gray_path = os.path.join("labels", rgb_file.replace(".png", ".png")) # 保持同名 rgb_to_gray_label(rgb_path, gray_path)逻辑说明:遥感标注工具(如 LabelMe、QGIS 插件)常导出 RGB 标签,但 PyTorch 的CrossEntropyLoss要求标签是[H,W]形状的 long tensor,值为 0~C-1。此脚本遍历每个像素,查表映射,避免cv2.cvtColor(..., cv2.COLOR_RGB2GRAY)把不同颜色压成同一灰度值。
2.2 遥感图像切块:不是等分,是带重叠的滑动窗口
遥感图太大(如 6000×7000),直接 resize 会损失细节。U-Net 输入通常为 256×256 或 512×512,但简单裁剪会导致边界信息丢失、地物被切断。必须用带重叠(overlap)的滑动窗口,且 overlap 要 ≥ stride 的 1/3,否则拼接时接缝明显。
# tile_images.py import numpy as np import tifffile as tiff import cv2 import os def tile_tif(image_path, label_path, tile_size=512, overlap=128): # 读取遥感图(tif 可能多波段,取前3或4波段) img = tiff.imread(image_path) # shape: (H, W, C) or (C, H, W) if img.ndim == 3 and img.shape[0] in [3, 4]: # CHW 格式 img = np.transpose(img, (1, 2, 0)) # HWC elif img.ndim == 2: # 单波段,扩维 img = np.expand_dims(img, axis=-1) label = cv2.imread(label_path, cv2.IMREAD_GRAYSCALE) # H,W h, w = img.shape[:2] stride = tile_size - overlap tiles_img, tiles_label = [], [] for i in range(0, h - tile_size + 1, stride): for j in range(0, w - tile_size + 1, stride): tile_img = img[i:i+tile_size, j:j+tile_size] tile_label = label[i:i+tile_size, j:j+tile_size] # 过滤掉全背景块(加速训练,减少噪声) if np.unique(tile_label).size > 1: # 至少含2类 tiles_img.append(tile_img) tiles_label.append(tile_label) return tiles_img, tiles_label # 示例:对 IMG_001.tif 切块 img_tiles, label_tiles = tile_tif( "data/images/IMG_001.tif", "data/labels/IMG_001.png", tile_size=512, overlap=128 ) # 保存切块 for idx, (t_img, t_lbl) in enumerate(zip(img_tiles, label_tiles)): cv2.imwrite(f"data/tiles/images/IMG_001_{idx:04d}.png", t_img) cv2.imwrite(f"data/tiles/labels/IMG_001_{idx:04d}.png", t_lbl)参数说明:
tile_size=512:U-Net 常用输入尺寸,显存友好(RTX 3060 可跑 batch_size=8)overlap=128:重叠 1/4,保证边缘地物完整,拼接时用加权平均融合(后文 predict.py 实现)np.unique(tile_label).size > 1:跳过纯背景块(如大片云层、黑边),减少无效训练
2.3 构建 PyTorch Dataset:处理多光谱归一化与数据增强
遥感图是多光谱,不能直接transforms.ToTensor()(它假设 RGB 三通道)。需手动归一化:对每个波段独立减均值除标准差。增强也需谨慎——遥感图旋转 90° 可能改变地物朝向(如道路变建筑),所以只做水平/垂直翻转、亮度对比度微调、高斯噪声。
# dataset.py import torch from torch.utils.data import Dataset import cv2 import numpy as np import os class RemoteSensingDataset(Dataset): def __init__(self, image_dir, label_dir, file_list, transform=None): self.image_dir = image_dir self.label_dir = label_dir self.file_list = file_list # ['IMG_001_0000', 'IMG_001_0001', ...] self.transform = transform # 遥感图各波段统计值(以 GF-2 数据为例,需按你数据集重算) self.mean = np.array([123.675, 116.28, 103.53, 129.1]) # B,G,R,NIR self.std = np.array([58.395, 57.12, 57.375, 56.8]) # 同上 def __len__(self): return len(self.file_list) def __getitem__(self, idx): name = self.file_list[idx] img_path = os.path.join(self.image_dir, f"{name}.png") lbl_path = os.path.join(self.label_dir, f"{name}.png") # 读取图像(支持 3/4 波段) img = cv2.imread(img_path, cv2.IMREAD_UNCHANGED) # H,W,C if img.ndim == 2: img = np.expand_dims(img, axis=-1) label = cv2.imread(lbl_path, cv2.IMREAD_GRAYSCALE) # 归一化:(x - mean) / std,逐波段 img = img.astype(np.float32) img = (img - self.mean[:img.shape[-1]]) / self.std[:img.shape[-1]] # 数据增强(仅翻转和噪声) if self.transform: if np.random.rand() > 0.5: img = np.fliplr(img) label = np.fliplr(label) if np.random.rand() > 0.5: img = np.flipud(img) label = np.flipud(label) if np.random.rand() > 0.7: noise = np.random.normal(0, 0.01, img.shape) img = np.clip(img + noise, -3, 3) # 限制范围 # 转 tensor img = torch.from_numpy(img.transpose(2, 0, 1)) # C,H,W label = torch.from_numpy(label).long() return img, label # 使用示例 train_files = [line.strip() for line in open("data/splits/train.txt")] train_dataset = RemoteSensingDataset( image_dir="data/tiles/images", label_dir="data/tiles/labels", file_list=train_files, transform=True )关键点说明:
self.mean/std必须按你数据集重算!用np.mean(img, axis=(0,1))计算各波段均值,否则归一化失效,模型不收敛。np.clip(img + noise, -3, 3):遥感图归一化后值域约 [-3,3],噪声不能超出,否则破坏分布。label.long():PyTorch 分类损失要求标签为 long 类型,否则报错expected scalar type Long but found Float。
3. U-Net 实现与训练:为什么原版结构要改这 3 处才能适配遥感
U-Net 原论文结构(2015)针对生物医学图像(512×512,单通道),直接搬来训遥感图会翻车:梯度爆炸、小目标漏检、训练震荡。我在 2022 年调试 GF-2 数据时,发现必须改这三处——不是炫技,是让模型在 30 个 epoch 内稳定收敛。
3.1 修改编码器:用 ResNet34 替代原版卷积块,解决梯度消失
原版 U-Net 编码器是 2×3×3 卷积 + ReLU,4 层下采样后特征图仅 32×32,小目标(如电线杆、小路)信息全丢。换成 ResNet34 作为编码器(即 Encoder),利用其残差连接和预训练权重,能保留更多空间细节。
# unet.py import torch import torch.nn as nn import torchvision.models as models class UNetEncoder(nn.Module): def __init__(self, pretrained=True): super().__init__() # 加载预训练 ResNet34,去掉最后两层(avgpool & fc) resnet = models.resnet34(pretrained=pretrained) self.first_conv = nn.Sequential( resnet.conv1, resnet.bn1, resnet.relu, resnet.maxpool ) self.layer1 = resnet.layer1 # 64 ch, 1/4 size self.layer2 = resnet.layer2 # 128 ch, 1/8 size self.layer3 = resnet.layer3 # 256 ch, 1/16 size self.layer4 = resnet.layer4 # 512 ch, 1/32 size def forward(self, x): x = self.first_conv(x) # [B,64,H/4,W/4] e1 = self.layer1(x) # [B,64,H/4,W/4] e2 = self.layer2(e1) # [B,128,H/8,W/8] e3 = self.layer3(e2) # [B,256,H/16,W/16] e4 = self.layer4(e3) # [B,512,H/32,W/32] return e1, e2, e3, e4 class UNetDecoder(nn.Module): def __init__(self, n_classes=5): super().__init__() # 上采样 + 跳连(注意通道数匹配) self.upconv4 = nn.ConvTranspose2d(512, 256, kernel_size=2, stride=2) self.decoder4 = self._make_block(512, 256) # 256(ch from up) + 256(ch from e3) self.upconv3 = nn.ConvTranspose2d(256, 128, kernel_size=2, stride=2) self.decoder3 = self._make_block(256, 128) self.upconv2 = nn.ConvTranspose2d(128, 64, kernel_size=2, stride=2) self.decoder2 = self._make_block(128, 64) self.upconv1 = nn.ConvTranspose2d(64, 32, kernel_size=2, stride=2) self.decoder1 = self._make_block(96, 32) # e1 是64ch,upconv1输出32ch,拼接后96ch self.final = nn.Conv2d(32, n_classes, kernel_size=1) def _make_block(self, in_ch, out_ch): return nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True), nn.Conv2d(out_ch, out_ch, 3, padding=1), nn.BatchNorm2d(out_ch), nn.ReLU(inplace=True) ) def forward(self, e1, e2, e3, e4): d4 = self.upconv4(e4) # [B,256,H/16,W/16] d4 = torch.cat([d4, e3], dim=1) # [B,512,H/16,W/16] d4 = self.decoder4(d4) # [B,256,H/16,W/16] d3 = self.upconv3(d4) # [B,128,H/8,W/8] d3 = torch.cat([d3, e2], dim=1) # [B,256,H/8,W/8] d3 = self.decoder3(d3) # [B,128,H/8,W/8] d2 = self.upconv2(d3) # [B,64,H/4,W/4] d2 = torch.cat([d2, e1], dim=1) # [B,128,H/4,W/4] d2 = self.decoder2(d2) # [B,64,H/4,W/4] d1 = self.upconv1(d2) # [B,32,H,W] d1 = self.decoder1(d1) # [B,32,H,W] return self.final(d1) # [B,5,H,W] class UNet(nn.Module): def __init__(self, n_classes=5, pretrained=True): super().__init__() self.encoder = UNetEncoder(pretrained=pretrained) self.decoder = UNetDecoder(n_classes=n_classes) def forward(self, x): e1, e2, e3, e4 = self.encoder(x) return self.decoder(e1, e2, e3, e4)为什么有效:ResNet34 的layer1输出特征图尺寸为H/4×W/4,比原版 U-Net 第一层H/2×W/2更细粒度,小目标定位更准;预训练权重(ImageNet)提供通用特征先验,缓解遥感小样本过拟合。
3.2 修改损失函数:Focal Loss + Dice Loss 混合,解决类别不平衡
遥感图中,水体可能占 5%,建筑占 30%,裸土占 65%。用CrossEntropyLoss会导致模型只优化大类,小类 IoU 接近 0。必须用Focal Loss(抑制易分样本) + Dice Loss(直接优化 IoU)混合:
# losses.py import torch import torch.nn as nn import torch.nn.functional as F class FocalLoss(nn.Module): def __init__(self, alpha=1, gamma=2, reduction='mean'): super().__init__() self.alpha = alpha self.gamma = gamma self.reduction = reduction def forward(self, inputs, targets): ce_loss = F.cross_entropy(inputs, targets, reduction='none') pt = torch.exp(-ce_loss) focal_weight = (1 - pt) ** self.gamma loss = focal_weight * ce_loss if self.reduction == 'mean': return loss.mean() return loss class DiceLoss(nn.Module): def __init__(self, smooth=1e-5): super().__init__() self.smooth = smooth def forward(self, logits, targets): probs = F.softmax(logits, dim=1) # [B,C,H,W] targets_one_hot = F.one_hot(targets, num_classes=logits.shape[1]) # [B,H,W,C] targets_one_hot = targets_one_hot.permute(0, 3, 1, 2).float() # [B,C,H,W] intersection = torch.sum(probs * targets_one_hot, dim=(2,3)) union = torch.sum(probs, dim=(2,3)) + torch.sum(targets_one_hot, dim=(2,3)) dice = (2. * intersection + self.smooth) / (union + self.smooth) return 1 - dice.mean() # 混合损失 class MixedLoss(nn.Module): def __init__(self, alpha=0.5): super().__init__() self.focal = FocalLoss(gamma=2) self.dice = DiceLoss() self.alpha = alpha # focal 权重 def forward(self, logits, targets): return self.alpha * self.focal(logits, targets) + (1 - self.alpha) * self.dice(logits, targets)参数说明:
alpha=0.5:平衡两项,实测在遥感数据上 0.4~0.6 效果稳定gamma=2:Focal Loss 标准值,加大难样本权重smooth=1e-5:Dice Loss 防止分母为 0
3.3 训练循环关键配置:学习率预热 + 余弦退火,避免初期震荡
U-Net 用预训练 ResNet 编码器,初期学习率太大易破坏特征提取能力。必须加warmup(前 5 个 epoch 线性增) + cosine annealing(后 25 个 epoch 平滑降):
# train.py import torch from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR, LinearLR def get_optimizer_and_scheduler(model, lr=1e-4, epochs=30): optimizer = AdamW(model.parameters(), lr=lr, weight_decay=1e-4) # warmup 5 epoch, then cosine decay to 1e-6 warmup_scheduler = LinearLR( optimizer, start_factor=0.01, end_factor=1.0, total_iters=5 ) main_scheduler = CosineAnnealingLR( optimizer, T_max=epochs-5, eta_min=1e-6 ) return optimizer, (warmup_scheduler, main_scheduler) # 训练主循环片段 model = UNet(n_classes=5).cuda() optimizer, schedulers = get_optimizer_and_scheduler(model, lr=3e-4, epochs=30) criterion = MixedLoss(alpha=0.5) for epoch in range(30): model.train() for img, lbl in train_loader: img, lbl = img.cuda(), lbl.cuda() pred = model(img) loss = criterion(pred, lbl) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) # 防梯度爆炸 optimizer.step() # 更新 scheduler:前5轮用 warmup,之后用 cosine if epoch < 5: schedulers[0].step() else: schedulers[1].step() # 验证 val_iou = validate(model, val_loader) print(f"Epoch {epoch}: Loss={loss.item():.4f}, Val IoU={val_iou:.4f}")血泪经验:不加clip_grad_norm_,第 3 轮就loss=nan;不用 warmup,前 10 轮 val IoU 在 0.1~0.3 间随机跳,根本看不到收敛趋势。
4. 避坑指南:遥感图像语义分割的 4 个致命错误与修复方案
这 4 个坑,我在毕设答辩现场见过至少 19 次——学生演示时模型输出一片红(全预测为建筑),导师皱眉问“为什么”,学生答“不知道,训练时 loss 降了”。其实全是数据或代码低级错误。按现象→原因→解决列清,照着检查 5 分钟就能救命。
4.1 现象:训练 loss 下降快,但验证 IoU 始终 < 0.2,且各类别混淆严重(如水体全标成道路)
原因:标签图是 RGB 彩色 PNG,但代码用cv2.imread(path)默认读为 BGR,再cv2.cvtColor(..., cv2.COLOR_BGR2GRAY)压成单通道,导致颜色映射错乱。例如 (0,0,255)(蓝)本该是水体 ID=3,但 BGR 读成 (255,0,0),映射成 ID=1(建筑)。
解决:
- 用
cv2.imread(path, cv2.IMREAD_GRAYSCALE)直接读灰度(前提是标签图已用 2.1 节脚本转好) - 或检查标签图是否真为灰度:
img = cv2.imread(path); print(img.shape),若为(H,W,3)则必错
4.2 现象:训练中途CUDA out of memory,即使 batch_size=1 也崩
原因:遥感图切块后仍存为.tif或未压缩.png,单张 512×512×4 波段图内存超 4MB,DataLoader 预加载全部数据到显存。
解决:
- 切块后统一转为
.png(有损压缩,体积降 60%):cv2.imwrite("xxx.png", tile_img, [cv2.IMWRITE_PNG_COMPRESSION, 3]) - DataLoader 加
pin_memory=False和num_workers=0(Windows 下多进程常引发内存泄漏)
4.3 现象:验证时 IoU=0.65,但导出的分割图全是“马赛克块”,边界锯齿、内部空洞
原因:预测时未做滑动窗口重叠融合,直接将整张大图 resize 到 512×512 再推理,再 resize 回原尺寸。resize 过程引入插值伪影,且小目标被模糊。
解决:必须用滑动窗口预测 + 重叠区域加权平均(见 5.2 节predict.py)
4.4 现象:模型在训练集上 IoU=0.8,验证集仅 0.4,过拟合严重
原因:数据增强过度。对遥感图做随机旋转(transforms.RandomRotation)会生成现实中不存在的地物朝向(如南北向道路变东西向),模型学到虚假特征。
解决:
- 删除所有旋转、仿射变换增强
- 仅保留:水平/垂直翻转、亮度/对比度 ±0.1、高斯噪声(σ=0.01)
- 加入CutMix(而非 MixUp):随机挖一个矩形块,用另一张图的对应块填充,保持地物完整性
注意:CutMix 实现需确保挖取区域不跨地物边界(用
cv2.findContours检测连通域,只在大块内部挖),否则会制造“鬼影”。
5. 模型部署与结果可视化:用 1 个 predict.py 脚本搞定全流程
毕设答辩核心是“让导师在自己电脑上看到效果”,不是跑出最高 IoU。本章教你写一个predict.py,输入一张.tif遥感图,输出带颜色覆盖的分割图(.png)和统计报表(.txt),双击即可运行,无需装环境。
5.1 predict.py 主逻辑:滑动窗口 + 重叠融合 + 颜色映射
# predict.py import torch import numpy as np import tifffile as tiff import cv2 import os from unet import UNet # 导入你训练好的模型 # 类别颜色映射(BGR 格式,OpenCV 用) CLASS_COLORS = { 0: (0, 0, 0), # 背景 - 黑 1: (0, 0, 255), # 建筑 - 红 2: (0, 255, 0), # 林地 - 绿 3: (255, 0, 0), # 水体 - 蓝 4: (0, 255, 255) # 道路 - 黄 } def load_model(model_path, n_classes=5): model = UNet(n_classes=n_classes) model.load_state_dict(torch.load(model_path, map_location='cpu')) model.eval() return model def sliding_predict(model, image, tile_size=512, overlap=128, device='cpu'): h, w = image.shape[:2] stride = tile_size - overlap # 初始化预测图和计数图(用于加权平均) pred_full = np.zeros((h, w), dtype=np.float32) count_full = np.zeros((h, w), dtype=np.int32) # 预计算重叠区域权重(中心高,边缘低) y_coords, x_coords = np.ogrid[:tile_size, :tile_size] center_y, center_x = tile_size // 2, tile_size // 2 dist = np.sqrt((y_coords - center_y)**2 + (x_coords - center_x)**2) weight_map = np.exp(-dist / (tile_size / 6)) # 高斯衰减 with torch.no_grad(): for i in range(0, h - tile_size + 1, stride): for j in range(0, w - tile_size + 1, stride): tile = image[i:i+tile_size, j:j+tile_size] # 归一化(同 dataset.py) tile = tile.astype(np.float32) mean = np.array([123.675, 116.28, 103.53, 129.1])[:tile.shape[-1]] std = np.array([58.395, 57.12, 57.375, 56.8])[:tile.shape[-1]] tile = (tile - mean) / std # 推理 tile_tensor = torch.from_numpy(tile.transpose(2,0,1)).unsqueeze(0).to(device) pred_tile = model(tile_tensor)[0] # [C,H,W] pred_prob = torch.softmax(pred_tile, dim=0).cpu().numpy() # [C,H,W] pred_class = np.argmax(pred_prob, axis=0) # [H,W] # 加权融合 pred_full[i:i+tile_size, j:j+tile_size] += pred_class * weight_map count_full[i:i+tile_size, j:j+tile_size] += weight_map # 取平均 pred_full = np.divide(pred_full, count_full, out=np.zeros_like(pred_full), where=count_full!=0) return np.round(pred_full).astype(np.uint8) def visualize_result(image, pred_mask, output_path): # 将预测结果映射为彩色图 h, w = pred_mask.shape color_mask = np.zeros((h, w, 3), dtype=np.uint8) for cls_id, color in CLASS_COLORS.items(): color_mask[pred_mask == cls_id] = color # 叠加原图(半透明) overlay = cv2.addWeighted(image[:, :, :3], 0.6, color_mask, 0.4, 0) cv2.imwrite(output_path, overlay) if __name__ == "__main__": # 配置 MODEL_PATH = "checkpoints/best_model.pth" INPUT_TIF = "data/images/IMG_001.tif" OUTPUT_DIR = "results/" os.makedirs(OUTPUT_DIR, exist_ok=True) # 加载模型 model = load_model(MODEL_PATH) model.to('cpu') # 为兼容导师电脑,强制 CPU 推理 # 读取遥感图 img = tiff.imread(INPUT_TIF) if img.ndim == 3 and img.shape[0] in [3,4]: img = np.transpose(img, (1,2,0)) # 滑动预测 print("Running sliding window prediction...") pred_mask = sliding_predict(model, img, tile_size=512, overlap=128) # 可视化 visualize_result(img, pred_mask, os.path.join(OUTPUT_DIR, "overlay.png")) # 生成统计报表 unique, counts = np.unique(pred_mask, return_counts=True) with open(os.path.join(OUTPUT_DIR, "stats.txt"), "w") as f: f.write("Class ID\tPixel Count\tPercentage\n") total = pred_mask.size for cls_id, cnt in zip(unique, counts): pct = cnt / total * 100 f.write(f"{cls_id}\t{cnt}\t{pct:.2f}%\n") print(f"Done! Results saved to {OUTPUT_DIR}")逻辑说明:
weight_map用高斯函数生成,中心权重 1.0,边缘渐降至 0.2,避免拼接线model.to('cpu'):确保导师电脑无 GPU 也能跑,RTX 3060 上 512×512 单块推理约 0.8s,整图(6000×
本文还有配套的精品资源,点击获取