據(jù)全流程詳解)
簡介這套7z壓縮包面向使用PyTorch處理高光譜圖像HSI的開發(fā)者與研究者針對高光譜數(shù)據(jù)通道多、內(nèi)存占用大、格式復雜等特點解決DataLoader加載數(shù)據(jù)時的讀取、預處理與批處理效率問題。包內(nèi)共7個文件含3個Python腳本數(shù)據(jù)加載、工具函數(shù)與訓練流程、2個pyc緩存文件以及2個mat格式的Indian Pines高光譜數(shù)據(jù)集整體壓縮后僅5.69MB腳本覆蓋dataloader封裝、Dataset構(gòu)造與訓練入口便于直接參照或改造。已有607人學習下載內(nèi)容聚焦從Dataset定義到collate_fn自定義的完整鏈路并涉及歸一化、多線程加載、pin_memory加速、shuffle隨機采樣與數(shù)據(jù)增強等實用設置。讀者可結(jié)合壓縮包內(nèi)的高光譜數(shù)據(jù)子目錄與數(shù)據(jù)集構(gòu)造模塊快速掌握針對高光譜多維數(shù)組的批處理方法和訓練腳本寫法同時理解緩存機制與自定義批處理函數(shù)對模型泛化和訓練效率的影響代碼結(jié)構(gòu)清晰、模塊劃分明確適合希望將PyTorch數(shù)據(jù)加載流程落地到HSI任務中的初中級工程師。 高光譜圖像這幾年在遙感深度學習里幾乎成了標配輸入——地物分類、變化檢測、異常目標識別動不動就是一個三維數(shù)據(jù)立方體直接喂給模型??珊芏嗳死@過了網(wǎng)絡結(jié)構(gòu)那關反而被最基礎的數(shù)據(jù)加載絆住高光譜數(shù)據(jù)不是一張圖而是一個長寬幾百、波段幾十上百的三維數(shù)組和torchvision里現(xiàn)成那套ImageFolder完全不是一個路子。這篇文章把我實際跑高光譜分類任務時怎么用PyTorch的DataLoader把數(shù)據(jù)真正“喂”進模型的全流程拆開講包括Dataset怎么寫、DataLoader參數(shù)怎么調(diào)、預處理怎么做、哪些坑我是踩過以后才明白的。內(nèi)容面向剛?cè)腴T遙感深度學習、或者被數(shù)據(jù)加載卡住的研究生和工程師有基礎代碼能力的人照著就能跑通。1. 高光譜數(shù)據(jù)加載為什么不能照搬普通圖像方案1.1 先搞清楚你手上到底是一份什么樣的數(shù)據(jù)高光譜遙感數(shù)據(jù)本質(zhì)上是一個三維數(shù)據(jù)立方體通常記作H×W×B。H和W是空間維度代表地物的長和寬B是光譜維度代表傳感器在連續(xù)電磁波譜上采樣的波段數(shù)。普通RGB圖像只有3個波段高光譜數(shù)據(jù)動輒上百個波段比如Indian Pines數(shù)據(jù)集是145×145像素、200個波段Pavia University是610×340像素、103個波段。每個像素不再是一個三通道的顏色值而是一條完整的光譜曲線這才是高光譜“識別地物”的核心價值所在。除了數(shù)據(jù)本身還有一個配套的標簽矩陣。以Indian Pines為例標簽是一個145×145的二維矩陣每個位置存一個類別編號比如0表示背景未標注1到16是不同地物類別。你的模型要做的事情就是根據(jù)中心像素周圍一個鄰域窗口內(nèi)的光譜和空間信息預測這個像素屬于哪一類。這個“鄰域窗口”的設定是理解高光譜DataLoader設計的關鍵起點。1.2 普通圖像加載方式在哪里行不通很多新手上來就嘗試torchvision的ImageFolder結(jié)果發(fā)現(xiàn)根本無從下手。原因是高光譜數(shù)據(jù)幾乎沒有現(xiàn)成的文件夾結(jié)構(gòu)一個.mat或者.h5文件里就裝著整幅圖像的全部數(shù)據(jù)標簽也不是文件名而是一個獨立的矩陣。更麻煩的是如果一個樣本就是一個完整的145×145×200的數(shù)據(jù)立方體顯存再大也塞不下——你總不能把一個整圖當成一個樣本去訓練吧。所以在高光譜分類里國內(nèi)的公開數(shù)據(jù)集Indian Pines、Pavia University、Salinas等普遍采用一個共同策略以每個像素為中心裁取一個固定大小的空間patch比如11×11或13×13把patch內(nèi)所有像素的所有波段作為輸入該中心像素的標簽作為輸出。這樣一來樣本數(shù)等于有效標注像素數(shù)每個樣本的尺寸是patch_size×patch_size×波段數(shù)既保留了空間上下文信息又把數(shù)據(jù)切成了適合訓練的塊。Dataset的核心工作就是把這個裁patch的過程封裝起來。2. 手寫自定義Dataset的核心實現(xiàn)2.1 Dataset接口只需要實現(xiàn)三個方法PyTorch定義Dataset類非常簡潔只要繼承torch.utils.data.Dataset然后實現(xiàn)__len__和__getitem__兩個方法就行。__len__返回樣本總數(shù)__getitem__給定一個索引返回一組訓練樣本和標簽。對于高光譜數(shù)據(jù)常見做法是在__init__階段把數(shù)據(jù)立方體和標簽矩陣讀進內(nèi)存同時把所有有效像素的坐標存成一個列表__getitem__里根據(jù)坐標索引去切patch。很多第一次寫Dataset的人會疑惑為什么不把patch提前切好存成數(shù)組原因很簡單——訓練時要做隨機采樣大部分數(shù)據(jù)集的標注像素有幾萬個Indian Pines約10249個像素有標注每個像素都要切一個11×11×200的patch提前切好意味著幾十G的內(nèi)存消耗和巨大的預處理時間完全不劃算。每次按需切片才是工程上合理的方式。2.2 兼容多格式數(shù)據(jù)讀取與歸一化高光譜公開數(shù)據(jù)集的存儲格式五花八門我遇到過的主要是三類MATLAB的.mat文件、HDF5的.h5文件、以及ENVI標準格式.hdr同名的.dat/.img文件。讀取邏輯建議在Dataset的__init__里做一次統(tǒng)一封裝這樣換數(shù)據(jù)集時只改讀取函數(shù)后續(xù)訓練邏輯完全不用動。用scipy.io的loadmat讀.mat用h5py讀.h5ENVI格式可以用spectral庫的envi.open接口。讀進來之后最重要的是歸一化。我自己的習慣是先做按波段的z-score標準化。做法是把數(shù)據(jù)從H×W×B reshape成(H*W)×B對每個波段算均值和標準差然后統(tǒng)一做標準化。這樣處理后每個波段的值都處在同一量級不會因為個別高反射波段在數(shù)值上主導梯度更新。注意標準差要加一個極小值比如1e-6防止某個全零的波段除零。2.3 完整可用的代碼模板下面這段代碼是我在實際項目里用的精簡版直接復制就能跑通Indian Pines這類數(shù)據(jù)。核心邏輯都在注釋我寫詳細一點。import numpy as np import torch from torch.utils.data import Dataset import scipy.io as sio class HyperspectralDataset(Dataset): def __init__(self, data_path, label_path, patch_size11, normalizationTrue, target_classNone): data_path: 高光譜數(shù)據(jù)文件(.mat或.h5) label_path: 標簽文件(.mat或.h5) patch_size: 空間鄰域窗口大小建議奇數(shù) # 讀取數(shù)據(jù)立方體shape: (H, W, B) if data_path.endswith(.mat): self.data sio.loadmat(data_path)[data].astype(np.float32) elif data_path.endswith(.h5): import h5py with h5py.File(data_path, r) as f: self.data f[data][:].astype(np.float32) else: raise ValueError(暫不支持該文件格式) # 讀取標簽矩陣shape: (H, W) if label_path.endswith(.mat): self.labels sio.loadmat(label_path)[label].astype(np.int64) else: import h5py with h5py.File(label_path, r) as f: self.labels f[label][:].astype(np.int64) h, w, b self.data.shape # 按波段z-score標準化 if normalization: flat self.data.reshape(-1, b) mean flat.mean(axis0) std flat.std(axis0) 1e-6 self.data ((self.data - mean) / std).astype(np.float32) self.patch_size patch_size self.pad patch_size // 2 # 對原圖做padding讓邊界像素也能切成完整patch self.data_padded np.pad(self.data, ((self.pad, self.pad), (self.pad, self.pad), (0, 0)), modereflect) self.labels_padded np.pad(self.labels, self.pad, modeconstant, constant_values0) # 收集所有有效像素的坐標 self.samples [] h_p, w_p self.labels_padded.shape for i in range(self.pad, h_p - self.pad): for j in range(self.pad, w_p - self.pad): if self.labels_padded[i, j] ! 0: self.samples.append((i, j)) def __len__(self): return len(self.samples) def __getitem__(self, idx): i, j self.samples[idx] # 切patchshape: (patch_size, patch_size, B) patch self.data_padded[i - self.pad : i self.pad 1, j - self.pad : j self.pad 1, :] # 轉(zhuǎn)成PyTorch需要的 (C, H, W) 格式 patch_tensor torch.from_numpy(patch.transpose(2, 0, 1)).float() label torch.tensor(self.labels_padded[i, j], dtypetorch.long) return patch_tensor, label這段代碼里有幾個細節(jié)值得說。第一我在__init__直接對原圖做了padding這樣邊界像素也能裁出完整的patch而且不用在__getitem__里做煩人的邊界判斷每次都是固定尺寸切片干凈很多。第二padding模式用reflect比constant填0更自然因為高光譜圖像相鄰像素光譜曲線本來就接近反射填充不會引入突兀的偽信息。第三標簽padding的地方填0這樣邊界區(qū)域即使被裁到也不會參與訓練因為0被我們當成背景過濾掉了。2.4 從整圖到patch的取舍邏輯為什么要用patch而不用單像素我最早做過一個對比實驗單像素輸入即1×1×B訓練出來的模型在Indian Pines上的總體精度大概比用11×11 patch低5到8個百分點。原因很直觀高光譜圖像里地物分類高度依賴空間紋理信息同一個光譜特征在農(nóng)田和城區(qū)可能代表完全不同的東西。patch相當于把中心像素周邊鄰居一起引入給模型提供了上下文。但patch也不是越大越好。patch過大有兩個問題一是類別邊界會被模糊邊緣像素的patch里混入了太多異類地物反而干擾分類二是計算量和顯存開銷隨patch面積平方增長。我實測下來Indian Pines用11×11或13×13比較均衡Pavia University空間分辨率相對高13×15左右的矩形patch也見過有人用。選patch時可以先固定一個值把流程跑通再去調(diào)參。3. DataLoader參數(shù)配置與性能細節(jié)3.1 batch_size和shuffle怎么設Dataset定義好了DataLoader就是個參數(shù)配置的事但參數(shù)配不好照樣出問題。先看batch_size。高光譜patch輸入是(B, C, H, W)的張量以11×11×200為例一個樣本的數(shù)據(jù)量是11×11×200×4字節(jié)約96KB看起來不大但batch累積起來就不一樣了。假設batch_size64一個batch的數(shù)據(jù)是64×96KB約6MB這只是輸入真正占顯存的是中間激活值模型越深、通道數(shù)越大顯存消耗越夸張。所以我的建議是先從batch_size16或32開始用nvidia-smi實時看顯存占用再逐步往上調(diào)找到一個“能跑滿GPU但不OOM”的值。shuffle參數(shù)在訓練集要設True這個大家基本都知道但要注意shuffle對高光譜數(shù)據(jù)的影響比普通圖像更大。高光譜數(shù)據(jù)集中同一個地物塊在空間上高度相關像素標簽是成片的。如果不shuffle一個batch里可能全是同一塊農(nóng)田的像素模型在這個batch里學到的全是局部特征loss曲線會像鋸齒一樣劇烈波動。shuffle之后每個batch都盡量混入不同類別的樣本訓練才穩(wěn)定。3.2 num_workers到底開多少num_workers控制DataLoader用幾個子進程來并行加載數(shù)據(jù)。對高光譜場景這里有個容易踩的大坑如果整個數(shù)據(jù)立方體都在內(nèi)存里每個worker進程會復制一份完整的數(shù)據(jù)副本。Indian Pines這種小數(shù)據(jù)量還好幾百MB撐死了但如果你處理的是航空影像拼接出來的大場景高光譜圖一個數(shù)據(jù)立方體可能好幾個GB開4個worker就意味著內(nèi)存直接翻4倍機器再大也容易扛不住。我的實際建議是先設num_workers0跑通確認邏輯沒問題后再嘗試增大。在Linux服務器上num_workers設為CPU核心數(shù)的一半通常性價比最高Windows環(huán)境下num_workers零點以上經(jīng)常報錯和系統(tǒng)多進程機制有關踩過這個坑之后我現(xiàn)在在Windows上干脆就一直用0。數(shù)據(jù)加載如果成了瓶頸優(yōu)先考慮用內(nèi)存映射或者提前把數(shù)據(jù)切成小塊而不是盲目加worker。3.3 pin_memory與數(shù)據(jù)類型轉(zhuǎn)換DataLoader里還有一個固定搭配建議直接加上pin_memoryTrue。這個參數(shù)的作用是把數(shù)據(jù)放進鎖頁內(nèi)存GPU訓練時從CPU傳到GPU可以走更快的數(shù)據(jù)通路幾乎是無本萬利的加速手段。唯一的代價是占用一點內(nèi)存對高光譜數(shù)據(jù)動輒幾百MB的數(shù)據(jù)集來說完全可以接受。數(shù)據(jù)類型方面要特別注意。我在__getitem__里返回的patch用torch.float32標簽用torch.long這是PyTorch訓練的標準配置。很多新手會忽略高光譜數(shù)據(jù)被讀進來時往往是float64比如從.mat讀出來默認就是double直接用float64的patch跑模型顯存直接翻倍速度還慢一半。所以在Dataset讀取階段一定要顯式.astype(np.float32)這個習慣能幫你少踩無數(shù)內(nèi)存坑。如果你用的是半精度混合精度訓練AMP那在訓練循環(huán)里做轉(zhuǎn)換就行Dataset里保持float32反而更靈活。4. 高光譜數(shù)據(jù)的預處理與數(shù)據(jù)增強4.1 歸一化是標配但歸一化的粒度有講究前面代碼里做了按波段的z-score標準化這在高光譜任務里幾乎是標配。但我看你數(shù)據(jù)的時候可以多做一步先把每一個波段的值統(tǒng)計一下分布高光譜數(shù)據(jù)經(jīng)常會遇到幾個波段全是噪聲或者全為零的情況比如水汽吸收波段這些波段如果直接參與訓練相當于往模型里灌垃圾信息。要么在預處理階段直接刪掉要么在做標準化時把方差極低的波段固定到一個小常數(shù)附近避免除零。歸一化粒度上有一個選擇全局歸一化還是按像素歸一化我傾向于按波段做全局標準化因為高光譜成像的物理含義是地表對太陽輻照的反射率不同波段的反射率有著不同的動態(tài)范圍統(tǒng)一到同一量級后模型學到的每個波段權(quán)重才有可比性。而按像素歸一化會破壞光譜間的相對關系反而不利于分類。4.2 光譜維度和空間維度的增強怎么做數(shù)據(jù)增強在高光譜任務里容易被忽略因為看起來“數(shù)據(jù)量挺大”——Indian Pines有幾萬像素標簽感覺足夠訓練了。但實際上很多地物類別樣本極少存在嚴重的類別不平衡。數(shù)據(jù)增強在高光譜里有一個獨特優(yōu)勢除了常規(guī)的空間增強翻轉(zhuǎn)、旋轉(zhuǎn)、隨機裁剪還能做光譜維度的增強這是普通RGB圖像做不到的。光譜增強里我試過兩種比較有效的方法。第一種是光譜加噪聲給patch的光譜維度加上服從高斯分布的小噪聲相當于模擬傳感器在不同光照條件下的噪聲變化。第二種是隨機波段丟棄每次訓練隨機丟掉5%到10%的波段逼模型學到冗余和魯棒的特征實測下來對提升泛化有穩(wěn)定幫助。空間增強方面翻轉(zhuǎn)和旋轉(zhuǎn)在patch級別操作即可但要注意驗證集和測試集不能做任何增強否則評價指標會虛高。4.3 類別不平衡問題的采樣策略高光譜數(shù)據(jù)集的類別不平衡非常嚴重比如Indian Pines里有些類別只有幾十個樣本而另一些有上千個。如果直接按原始分布訓練模型會學成“多數(shù)類主導”少數(shù)類幾乎預測不出來。這時候可以給DataLoader配一個WeightedRandomSampler權(quán)重和每個類別的樣本數(shù)成反比讓稀有類別在采集時獲得更高的概率。具體做法是統(tǒng)計每個類別的像素數(shù)量計算權(quán)重數(shù)組傳給采樣器。不過這里有個現(xiàn)實問題加權(quán)采樣可能會導致多數(shù)類欠擬合總體精度反而下降。我自己的經(jīng)驗是如果目標是論文里的Overall Accuracy對比可以先留著不平衡不做處理把基線跑出來如果目標是實際應用中的地物識別那加權(quán)采樣或者Focal Loss值得優(yōu)先嘗試。作為一個工程問題先把加權(quán)采樣器實現(xiàn)了對比一下再定。5. 實操高頻問題與排查經(jīng)驗5.1 加載速度慢訓練一直在等數(shù)據(jù)如果訓練時GPU利用率經(jīng)常掉到50%以下很大概率是數(shù)據(jù)加載成了瓶頸。高光譜數(shù)據(jù)計算量本身不大瓶頸往往在磁盤IO和內(nèi)存拷貝上。第一步可以檢查你讀入的是不是壓縮格式比如.mat里默認可能用了壓縮存儲每次讀取都要解壓這個我在實踐中遇到多次如果是解決思路是預先轉(zhuǎn)成內(nèi)存友好的.npy格式。第二步檢查__getitem__里有沒有做了多余的計算比如每次都在里面重新做切片、標準化等重復運算。標準化應該提前在__init__完成__getitem__只負責最輕量的切片和類型轉(zhuǎn)換。5.2 內(nèi)存暴漲程序直接被殺內(nèi)存問題最常見的原因就是前面提到的多進程復制。另外還有一類情況容易被忽略你把整個數(shù)據(jù)集在Dataset里讀了一遍但在預處理時又用np.concatenate或者Python列表不斷追加導致多份拷貝同時存在。老話重提高光譜數(shù)據(jù)處理最好全程用Numpy數(shù)組減少不必要的拷貝。如果數(shù)據(jù)真的太大可以選擇在__getitem__里按需讀取HDF5文件的特定區(qū)域HDF5天然支持部分讀取比一次性加載整個大文件更優(yōu)雅。5.3 驗證和測試階段的分割要小心訓練集和驗證集的劃分很多人直接在像素級別上隨機劃分這在遙感場景會有嚴重問題同一個地物的相鄰像素高度相關隨機劃分會把“劇透”信息泄漏進驗證集導致驗證精度虛高。更重要的是如果訓練集和驗證集有大量空間重疊你評估的不是泛化能力而是記憶能力。正確的做法是按空間區(qū)域劃分或者對每個類別按像素列表分層抽樣但保證同一類別的訓練和驗證像素盡量遠離。工業(yè)界還有一種做法是分塊留出比如整圖按網(wǎng)格切塊把一部分塊整體作為驗證集這樣更貼近真實應用場景。另外有一個小坑Dataset的__getitem__每次返回的patch都是獨立切片驗證時需要逐patch推理再拼接成完整預測圖。這個過程注意也要padding一致否則拼接出來的預測圖邊緣會對不齊。5.4 一個完整的訓練調(diào)用示例最后把DataLoader部分整合起來方便你直接參考標準用法。from torch.utils.data import DataLoader from torch.utils.data.sampler import WeightedRandomSampler import numpy as np # 實例化Dataset train_ds HyperspectralDataset( data_pathIndian_Pines.mat, label_pathIndian_Pines_gt.mat, patch_size11, normalizationTrue ) # 按類別數(shù)量計算采樣權(quán)重 labels train_ds.labels # (H, W) unique, counts np.unique(labels[labels ! 0], return_countsTrue) class_count dict(zip(unique, counts)) sample_weights [] for i, j in train_ds.samples: cls train_ds.labels[i - train_ds.pad, j - train_ds.pad] sample_weights.append(1.0 / class_count[cls]) sampler WeightedRandomSampler(sample_weights, num_sampleslen(sample_weights), replacementTrue) train_loader DataLoader( train_ds, batch_size32, shuffleFalse, # 使用sampler時必須置False samplersampler, # 可替換為None來關閉加權(quán)采樣 num_workers4, pin_memoryTrue, drop_lastTrue ) for batch_idx, (patches, targets) in enumerate(train_loader): patches patches.cuda() targets targets.cuda() # 這里就是你的模型前向和反向代碼這個加載環(huán)節(jié)跑通之后剩下的模型結(jié)構(gòu)、損失函數(shù)、評估指標都和水到渠成一樣。但環(huán)境依賴那里多說一句PyTorch和CUDA版本的匹配是個老生常談的問題我第一次裝GPU版的時候就被版本不兼容坑過一整天直接按照官方提供的組合命令來裝盡量別混裝。我在實際做高光譜分類項目的過程中反復調(diào)整最多的不是網(wǎng)絡層數(shù)反而是數(shù)據(jù)加載和預處理這部分。尤其當你在多個數(shù)據(jù)集上做對比實驗時Dataset寫得好不好直接決定你后續(xù)的工作量。把數(shù)據(jù)讀取、歸一化、patch采樣、加載調(diào)度這四件事固化成一個通用模塊以后換任何高光譜數(shù)據(jù)集都能幾分鐘內(nèi)適配這件事值得你花時間一次性做扎實。本文還有配套的精品資源點擊獲取