简介针对高光谱图像在深度学习训练中数据维度高、读取慢、预处理复杂等痛点这套资料提供了基于PyTorch DataLoader的完整数据加载示例适合使用Python与PyTorch进行高光谱图像分类、地物识别等任务的研究者或初学者参考。压缩包共7个文件主要包含3个Python脚本涵盖数据集定义、工具函数与训练入口另有2个编译生成的pyc文件以及2个MAT格式的高光谱数据文件如IndianPines标准数据集整体仅5.69MB轻量易用。目前已有609人浏览学习。示例代码围绕自定义Dataset类展开从数据读取、预处理到批处理多个层次逐步演示清晰覆盖了长度与获取方法的实现、光谱通道归一化、批大小与工作线程配置、随机采样、内存加速以及自定义collate_fn处理多通道样本等核心操作并附有完整训练脚本可直接运行查看加载与训练效果。这套资料既适合希望快速上手高光谱数据加载的初学者也可作为中高级开发者的参考工具帮助规避内存溢出、加载缓慢、批量合并出错等常见问题。1. 高光谱数据用DataLoader为什么不能照搬RGB那套先算清三笔账高光谱数据和普通图像的分量完全不同。一张200波段的影像空间尺寸512×512float32存储整图约200MBICVL这类.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_idxNone): with h5py.File(file_path, r) as f: data f[key] if band_idx is None: patch data[row0:row0rows, col0:col0cols, :] else: patch data[row0:row0rows, col0:col0cols, 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_offcol0, row_offrow0, widthcols, heightrows) patch src.read(windowwin) # 返回形状 (band, rows, cols) return patchrasterio的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, source2, destination0) 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, keyradiance, label_keymap, patch_size11, meanNone, stdNone): 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)), modereflect) 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.wx 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到65535radiance数据范围更大。常见错误是每个patch单独做归一化这会严重破坏数据分布——同一个地物在不同位置亮度不同每patch独立归一化等于把亮度差异全抹掉模型学到的不是真实光谱特性。正确做法是在Dataset构造函数里提前算好整图的mean和stdgetitem里统一套用。但整图统计如果直接读全图内存会爆所以分块统计import h5py import numpy as np def compute_global_stats(mat_path, keyradiance, block_rows64): with h5py.File(mat_path, r) as f: data f[key] h, w, b data.shape mean np.zeros(b, dtypenp.float64) sq_mean np.zeros(b, dtypenp.float64) count 0 for i in range(0, h, block_rows): block data[i:iblock_rows, :, :].reshape(-1, b) mean block.sum(axis0) sq_mean (block.astype(np.float64) ** 2).sum(axis0) 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单样本才135KBbatch_size64也就8.6MB但这只是输入。真正占显存大头的是中间特征图和反向传播的梯度一个简单的两层3D卷积网络中间特征往往比输入大几十倍。经验做法是先从batch_size16或32起步跑一个epoch观察显存占用再往上加。不同patch尺寸的参考patch_size波段数单样本大小建议初始batch_size7×710019KB6411×1120095KB3213×13200135KB1619×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_size16, shuffleTrue, num_workers4, pin_memoryTrue, prefetch_factor4, persistent_workersTrue, )参数说明num_workers4表示开4个子进程做数据加载能并行读磁盘pin_memoryTrue是在GPU训练时锁页内存减少host到device的拷贝时间prefetch_factor4表示每个worker提前预取4个batchpersistent_workersTrue让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, dim0) return patches, torch.as_tensor(labels, dtypetorch.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×20033800个值里只有一行13×200被替换掉。大部分读盘和moveaxis计算其实在做重复功。可以用lru_cache给getitem加一层缓存把最近读过的patch存下来from functools import lru_cache class CachedHSIDataset(HSI_PatchDataset): lru_cache(maxsize4096) 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)), modereflect) 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。maxsize4096在高光谱patch下大概占几百MB内存具体根据服务器内存调整。注意num_workers0时每个worker进程有独立的一份缓存不会互相共享所以内存上限要按worker数估算。另一个提速技巧是给验证集单独写一个Sampler不要shuffle而是等步长采样每隔N个像素取一个样本。高光谱相邻像素几乎相同全量验证有一大半是重复计算采样子集后验证速度和精度都能兼顾。这个Sampler直接传给DataLoader的sampler参数即可。我自己最开始做高光谱分类时也是先整图load然后手动切patch内存翻车了两三次才换成h5py懒加载后来又栽在归一化上直到把全局统计改到构造函数里loss曲线才恢复正常。这套加载方案跑过ICVL也跑过自己拍的无人机高光谱数据稳定性和速度都让人放心希望帮到你。本文还有配套的精品资源点击获取