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