☰
高光谱数据加载避坑指南:PyTorch DataLoader与Dataset优化实践
2026/9/25 6:26:40 网站建设 项目流程

简介:针对高光谱图像在深度学习训练中数据维度高、读取慢、预处理复杂等痛点,这套资料提供了基于PyTorch DataLoader的完整数据加载示例,适合使用Python与PyTorch进行高光谱图像分类、地物识别等任务的研究者或初学者参考。压缩包共7个文件,主要包含3个Python脚本,涵盖数据集定义、工具函数与训练入口,另有2个编译生成的pyc文件以及2个MAT格式的高光谱数据文件(如IndianPines标准数据集),整体仅5.69MB,轻量易用。目前已有609人浏览学习。示例代码围绕自定义Dataset类展开,从数据读取、预处理到批处理多个层次逐步演示,清晰覆盖了长度与获取方法的实现、光谱通道归一化、批大小与工作线程配置、随机采样、内存加速以及自定义collate_fn处理多通道样本等核心操作,并附有完整训练脚本,可直接运行查看加载与训练效果。这套资料既适合希望快速上手高光谱数据加载的初学者,也可作为中高级开发者的参考工具,帮助规避内存溢出、加载缓慢、批量合并出错等常见问题。

1. 高光谱数据用DataLoader为什么不能照搬RGB那套:先算清三笔账

高光谱数据和普通图像的分量完全不同。一张200波段的影像,空间尺寸512×512,float32存储,整图约200MB,ICVL这类.mat数据集一个场景动辄上GB。PyTorch的DataLoader默认每个epoch从磁盘重新读数据、在getitem里做预处理,如果第一步就把整幅影像一次性读进内存,训练还没开始机器先卡死;更隐蔽的是,高光谱影像通常是(H, W, B)布局,波段在最后一位,而PyTorch卷积要的是(B, C, H, W),直接送进去必然踩维度坑。这篇笔记围绕PyTorch DataLoader展开,把高光谱文件格式、Dataset编写、batch参数取舍、归一化和采样策略说清楚,适合正在做高光谱地物分类、目标探测、波段筛选,且被数据加载拖住训练进度的同行。文章里给的代码都是可以直接抄的短方案,注释里说明为什么这么写。

2. 从原始文件到Tensor:按格式选路线,按内存定策略

2.1 先认清数据文件是哪一种:.mat、GeoTIFF、ENVI的识别方法

高光谱数据流通得最多的是三种形态。第一种是.mat格式,ICVL高光谱数据集就是这种,注意MATLAB高版本保存的.mat本质是HDF5容器,h5py可以直接打开,不需要装MATLAB。第二种是GeoTIFF,一个文件带N个波段,栅格软件常见,rasterio按波段号读取。第三种是ENVI的.dat加.hdr,科研和商业软件最常用,光谱库和ENVI用户手里一把一把的。识别方法很简单,先打印元信息再动手写加载逻辑:

import h5py f = h5py.File('./icvl_grn_500.mat', 'r') print(list(f.keys())) # 高光谱场景数据一般叫 'radiance' 或 'reflectance',以实际打印为准 data = f['radiance'] print(data.shape, data.dtype)
gdalinfo scene.dat | grep -E 'Type=|Size is|Band Count'

第一段代码打开.mat文件并列出内部所有key。为什么强调打印key,因为.mat的key命名没有统一规范,同一个数据集有人存成radiance,有人存成reflectance,还有人存成data,写死名字在别人的数据上必翻车。第二段gdalinfo是GDAL工具自带的命令,看ENVI格式的波段数、数据类型、尺寸非常快,比写Python脚本省事。

2.2 不要让整幅图进内存:h5py和rasterio的懒加载切片

高光谱加载真正考验人的是内存策略。h5py打开文件返回的是文件对象,不会立刻把数据搬进RAM,只有data[100:200, 100:200, :]这种切片才触发实际磁盘读取。rasterio同理,read时传window参数只读感兴趣区域。写一个按坐标取patch的函数:

import h5py import numpy as np def load_patch_mat(file_path, key, row0, col0, rows, cols, band_idx=None): with h5py.File(file_path, 'r') as f: data = f[key] if band_idx is None: patch = data[row0:row0+rows, col0:col0+cols, :] else: patch = data[row0:row0+rows, col0:col0+cols, band_idx] return np.asarray(patch)

参数说明:row0、col0是patch左上角坐标,rows、cols是patch的高度和宽度;band_idx传None表示取全部波段,传入整数list就可以做波段选择。注意函数里用with块打开文件,函数结束时h5py释放文件句柄,但返回的np.array已经拷贝到内存里,和文件句柄无关,可以放心用。rasterio版写法略有不同:

from rasterio.windows import Window import rasterio def load_patch_tif(file_path, row0, col0, rows, cols): with rasterio.open(file_path) as src: win = Window(col_off=col0, row_off=row0, width=cols, height=rows) patch = src.read(window=win) # 返回形状 (band, rows, cols) return patch

rasterio的read返回的是(C, H, W),波段在第一位,和h5py切出来的(H, W, B)不一样。这个小差别后面在Dataset里处理维度转换时要格外小心,两种函数不能混用同一个transpose逻辑。

2.3 维度约定:把(H,W,B)转成卷积认识的(B,C,H,W)

PyTorch的Conv2d输入通道在第二个维度,而高光谱影像习惯上把波段放最后。直接从文件切出来的patch是(H, W, B),在Dataset里需要变成(C, H, W):

patch = np.moveaxis(patch, source=2, destination=0) patch = np.ascontiguousarray(patch) # transpose后内存布局不连续,必须整理

moveaxis做的是视图变换,返回的数组在内存里是stride交错排列的,直接torch.from_numpy再接卷积会触发隐式拷贝,每次getitem都复制一次,训练速度慢得明显。用ascontiguousarray把内存布局固化。高光谱波段少则几十、多则几百,通道维放前面之后,接3D卷积、光谱注意力、波段注意力模块都不用手忙脚乱。这个顺序约定越早统一到Dataset里,后面的模型代码越干净。

3. 手写高光谱Dataset:让每个getitem只读一小块

3.1 最小可运行的逐像素PatchDataset

高光谱分类最常见的任务是以每个像素为中心裁patch,一个样本就是一小块影像。写Dataset时要抛弃“先加载整图再索引”的思维,正确姿势是构造函数只保存文件路径和坐标换算信息,真正读磁盘的动作放getitem内部:

from torch.utils.data import Dataset import h5py import numpy as np import torch class HSI_PatchDataset(Dataset): def __init__(self, mat_path, key='radiance', label_key='map', patch_size=11, mean=None, std=None): self.file = h5py.File(mat_path, 'r') # 保持打开,getitem里切片 self.data = self.file[key] self.labels = self.file[label_key] self.p = patch_size self.half = patch_size // 2 self.h, self.w = self.data.shape[:2] self.bands = self.data.shape[2] self.mean = mean self.std = std def __len__(self): return self.h * self.w def __getitem__(self, idx): y = idx // self.w x = idx % self.w y0 = max(0, y - self.half) y1 = min(self.h, y + self.half + 1) x0 = max(0, x - self.half) x1 = min(self.w, x + self.half + 1) patch = self.data[y0:y1, x0:x1, :] patch = np.pad( patch, ((y0 - (y - self.half), (y + self.half + 1) - y1), (x0 - (x - self.half), (x + self.half + 1) - x1), (0, 0)), mode='reflect') patch = np.moveaxis(patch, 2, 0) patch = torch.from_numpy(np.ascontiguousarray(patch)).float() if self.mean is not None: patch = (patch - self.mean) / self.std label = int(self.labels[y, x]) return patch, label

这段代码可以直接跑,注意key名称要根据实际数据修改。__len__返回像素总数,getitem把一维索引换算成行列坐标;h5py切片是内存映射读取,不读整图;边界像素裁不满patch_size时,用np.pad以reflect模式补齐。对遥感影像边缘,reflect比零填充和常数填充合理,不会在边缘引入虚假的黑色边框。label的取法根据实际数据调整,语义分割数据一般是和影像同尺寸的标签图。

3.2 坐标换算和边界填充的细节:为什么用reflect而不是zero

逐像素分类的Dataset,本质上就是“把一个像素位置映射到一个小patch”。换算公式很简单:y = idx // self.w,x = idx % self.w,一维索引转二维坐标。但边界像素的采样有个坑:图像边缘不够patch_size时,切片尺寸会比预期小。上面代码用pad参数动态补齐,把“实际切出来的区域相对标准patch的偏移量算出来,然后左右上下各补多少”一次算清楚。有个更省事的写法是image自己先pad再切,但那样浪费时间也浪费内存,因为整图pad要复制一整份数据。高光谱一张图几百MB,整图pad非常不划算,动态算偏移量才是正确做法。reflect模式会把边缘像素倒影过来,例如一行像素[1,2,3,4],在左边补两个值就是[2,1,1,2],在光谱上这种延续比补0更接近真实地物分布。

3.3 归一化参数放哪:统计整图mean/std,不要在getitem里现算

高光谱数据的取值范围因传感器而异,反射率数据有的在0到1,有的在0到65535,radiance数据范围更大。常见错误是每个patch单独做归一化,这会严重破坏数据分布——同一个地物在不同位置亮度不同,每patch独立归一化等于把亮度差异全抹掉,模型学到的不是真实光谱特性。正确做法是在Dataset构造函数里提前算好整图的mean和std,getitem里统一套用。但整图统计如果直接读全图内存会爆,所以分块统计:

import h5py import numpy as np def compute_global_stats(mat_path, key='radiance', block_rows=64): with h5py.File(mat_path, 'r') as f: data = f[key] h, w, b = data.shape mean = np.zeros(b, dtype=np.float64) sq_mean = np.zeros(b, dtype=np.float64) count = 0 for i in range(0, h, block_rows): block = data[i:i+block_rows, :, :].reshape(-1, b) mean += block.sum(axis=0) sq_mean += (block.astype(np.float64) ** 2).sum(axis=0) count += block.shape[0] mean /= count sq_mean /= count std = np.sqrt(np.maximum(sq_mean - mean ** 2, 0)) return mean.astype(np.float32), std.astype(np.float32)

分块统计的思路是把图像按行切成一段段,每段只有64×w×b这么大,内存占用可控。累加器用float64,防止200波段、几十万像素的float数据累加时精度漂移。sq_mean - mean ** 2在浮点运算下可能因为舍入出现极小负数,外面套一层np.maximum兜底,避免开根号时得到nan。

4. DataLoader参数配置:batch_size、num_workers、shuffle、collate_fn一次配齐

4.1 按显存预算反推batch_size

高光谱patch看起来很小,13×13×200波段,float32单样本才135KB,batch_size=64也就8.6MB,但这只是输入。真正占显存大头的是中间特征图和反向传播的梯度,一个简单的两层3D卷积网络,中间特征往往比输入大几十倍。经验做法是先从batch_size=16或32起步跑一个epoch,观察显存占用再往上加。不同patch尺寸的参考:

patch_size波段数单样本大小建议初始batch_size
7×710019KB64
11×1120095KB32
13×13200135KB16
19×19200289KB8

这个表是按float32输入的保守值,如果网络里用了大kernel的3D卷积或transformer结构,batch_size还要再减半。显存不够时优先减batch_size而不是patch_size,因为patch太小会丢掉空间上下文信息,对高光谱分类精度影响明显。

4.2 num_workers怎么调:高光谱加载是IO密集,不是算力密集

高光谱的getitem主要在做磁盘读取和numpy切片,这是IO密集操作。num_workers默认是0,加载在主进程执行,GPU算得快时数据供不上,训练曲线会出现明显的“停顿”。建议从4开始试,8封顶。DataLoader的标准配置:

from torch.utils.data import DataLoader train_loader = DataLoader( train_dataset, batch_size=16, shuffle=True, num_workers=4, pin_memory=True, prefetch_factor=4, persistent_workers=True, )

参数说明:num_workers=4表示开4个子进程做数据加载,能并行读磁盘;pin_memory=True是在GPU训练时锁页内存,减少host到device的拷贝时间;prefetch_factor=4表示每个worker提前预取4个batch;persistent_workers=True让worker进程跨epoch复用,避免每个epoch结束后重新fork子进程的开销。注意Windows下多进程DataLoader要求主训练代码放在if __name__ == '__main__':保护块里,否则子进程递归执行入口直接崩溃,这是Windows多进程模型的硬性限制。

4.3 collate_fn要不要写:固定patch尺寸可以不写,多尺度训练必须写

默认collate_fn会把list里的样本自动stack成tensor,前提是每个样本形状完全一致。固定patch_size的高光谱Dataset不需要自定义collate_fn。做多尺度训练或数据增强后patch尺寸不一致时,默认collate会直接报错,需要自己写:

def collate_hsi(batch): patches, labels = zip(*batch) patches = torch.stack(patches, dim=0) return patches, torch.as_tensor(labels, dtype=torch.long)

这段代码的作用是把batch里的patch沿第0维堆叠成四维tensor,标签转成long类型,适配CrossEntropyLoss。缺点是这样要求所有patch形状一致,做多尺度增强时还得配合自适应池化或者把patch变成统一尺寸。

5. 高光谱DataLoader避坑:四条翻车经历和对应修法

5.1 翻车:训练进程没开始,内存先被整个数据集撑爆

现象:脚本一启动内存占用直接飙到90%以上,还没开始训练就卡死,甚至整个服务器无响应。

原因:Dataset的构造函数里写了self.data = self.file[key][...],带了省略号等价于把整个数据集读进内存。高光谱文件动辄几个GB,一段代码就毁掉整台机器。

解决:h5py文件对象保留在Dataset里,getitem里再切片。检查方法很简单,打开任务管理器或top,看内存上涨发生在数据加载阶段还是训练阶段。出现整图加载时,把构造函数里所有带[...]的赋值都拆掉,只保留文件对象和shape信息。

5.2 翻车:每个patch单独做归一化,loss下降慢且验证集震荡

现象:训练loss下降非常缓慢,验证loss曲线像锯齿一样上下跳,模型精度始终上不去。

原因:每个patch单独做min-max归一化,不同patch的亮度和对比度被强行拉到一样的范围。高光谱影像里阴影区、亮目标、水体相对反射率本身就不同,独立归一化等于把地物间的真实亮度差异抹掉了,模型第一层学到的特征不再一致。

解决:用整图统计的全局mean和std归一化,把统计结果存到Dataset的成员变量里,getitem里只做减法和除法。修改后loss曲线明显平滑,收敛速度也快得多。

5.3 翻车:训练集验证集随机切分,精度虚高到不敢信

现象:随机切分时训练精度98%,验证精度99%,模型下放到新场景精度掉到70%。

原因:高光谱相邻像素空间相关性强,同一地物区域内的像素几乎一样。随机切分会让训练集和验证集出现大量重叠区域像素,这就是数据泄漏,验证精度被严重高估。

解决:按空间区块切分。把影像划分成互不相交的若干大区块,训练集和验证集各自用完整的区块,保证验证集里的像素在空间上完全隔离。切换后验证精度会明显下降,但这个数字才是真实水平。

5.4 翻车:.mat转tensor后通道维错乱,卷积直接报维度错误

现象:RuntimeError: Expected 4D input [N, C, H, W], got [N, H, W, C],或者模型跑起来特别慢。

原因:h5py切出来的patch是(H, W, B),PyTorch期望(B, C, H, W)。波段维位置不对,卷积层会报错或隐式做低效的通道搬移。

解决:在Dataset里统一做np.moveaxis(patch, 2, 0)加np.ascontiguousarray(patch),这是必须的一步,不是优化项。调试时可以先跑一个batch打印shape,确认四维分别是batch、波段、高、宽。

5.5 翻车:拖着radiance当reflectance用,模型难以跨数据集迁移

现象:在ICVL上训好的模型,换一个数据集效果大跌;同一个场景不同时间拍摄的影像,预测结果差异巨大。

原因:radiance是传感器接收的辐亮度,受光照、大气、观测角度影响;reflectance是地物本身的光谱反射特性。两者之间差一个大气校正,直接混用等于让模型同时学习“地物是什么”和“当时天气怎么样”两件事。

解决:在Dataset入口统一做转反射率处理。高光谱如何转反射率,常见做法是用ENVI做大气校正,或者对已知参考白板做经验定标。至少要在代码里明确标注当前数据集是radiance还是reflectance,训练和验证用同一类数据。

6. 进阶:给DataLoader做一层小缓存,把相同patch的重复计算消掉

高光谱逐像素采样有个天然特点:相邻像素的patch高度重叠,13×13的patch中心每移动一个像素,13×13×200=33800个值里只有一行13×200被替换掉。大部分读盘和moveaxis计算其实在做重复功。可以用lru_cache给getitem加一层缓存,把最近读过的patch存下来:

from functools import lru_cache class CachedHSIDataset(HSI_PatchDataset): @lru_cache(maxsize=4096) def read_raw_patch(self, y, x): y0 = max(0, y - self.half) y1 = min(self.h, y + self.half + 1) x0 = max(0, x - self.half) x1 = min(self.w, x + self.half + 1) patch = self.data[y0:y1, x0:x1, :] patch = np.pad( patch, ((y0 - (y - self.half), (y + self.half + 1) - y1), (x0 - (x - self.half), (x + self.half + 1) - x1), (0, 0)), mode='reflect') return np.moveaxis(patch, 2, 0) def __getitem__(self, idx): y = idx // self.w x = idx % self.w patch = torch.from_numpy(self.read_raw_patch(y, x)).float() if self.mean is not None: patch = (patch - self.mean) / self.std return patch, int(self.labels[y, x])

lru_cache的作用是把读patch和维度转换缓存起来,key是(y, x),后续碰到相同坐标直接返回缓存结果,不再走磁盘读取和reflect pad。maxsize=4096在高光谱patch下大概占几百MB内存,具体根据服务器内存调整。注意num_workers>0时每个worker进程有独立的一份缓存,不会互相共享,所以内存上限要按worker数估算。

另一个提速技巧是给验证集单独写一个Sampler,不要shuffle而是等步长采样,每隔N个像素取一个样本。高光谱相邻像素几乎相同,全量验证有一大半是重复计算,采样子集后验证速度和精度都能兼顾。这个Sampler直接传给DataLoader的sampler参数即可。

我自己最开始做高光谱分类时,也是先整图load然后手动切patch,内存翻车了两三次才换成h5py懒加载;后来又栽在归一化上,直到把全局统计改到构造函数里,loss曲线才恢复正常。这套加载方案跑过ICVL也跑过自己拍的无人机高光谱数据,稳定性和速度都让人放心,希望帮到你。

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

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

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

立即咨询