據(jù)集實(shí)戰(zhàn)指南:InMemoryDataset 與 Dataset 的創(chuàng)建、加載與擴(kuò)展)
PyTorch Geometric 自定義圖數(shù)據(jù)集實(shí)戰(zhàn)指南InMemoryDataset 與 Dataset 的創(chuàng)建、加載與擴(kuò)展【免費(fèi)下載鏈接】pytorch_geometricGraph Neural Network Library for PyTorch項(xiàng)目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric雖然 PyTorch GeometricPyG已經(jīng)內(nèi)置了大量高質(zhì)量數(shù)據(jù)集如 Planetoid、QM9、OGB 系列等見 torch_geometric/datasets但當(dāng)你面對(duì)自采數(shù)據(jù)或非公開數(shù)據(jù)時(shí)往往需要親手實(shí)現(xiàn)自己的數(shù)據(jù)集類。本指南以官方文檔 docs/source/tutorial/create_dataset.rst 為核心結(jié)合倉庫源碼系統(tǒng)講解 PyG 數(shù)據(jù)集的兩大抽象基類InMemoryDataset與Dataset的設(shè)計(jì)思想、目錄約定、四個(gè)核心鉤子方法與完整實(shí)現(xiàn)范例并深入剖析collate/save/load的底層機(jī)制。讀完本文你將能夠獨(dú)立編寫可復(fù)現(xiàn)、可緩存、可增量加載的自定義圖數(shù)據(jù)集并與DataLoader無縫銜接。數(shù)據(jù)集抽象基類兩類選擇一種約定PyG 在torch_geometric.data命名空間中提供了兩個(gè)抽象基類參見 torch_geometric/data/init.py 的導(dǎo)出torch_geometric.data.Dataset通用數(shù)據(jù)集基類繼承自torch.utils.data.Dataset適用于無法整體裝入內(nèi)存的大規(guī)模數(shù)據(jù)集torch_geometric.data.InMemoryDataset繼承自Dataset當(dāng)整個(gè)數(shù)據(jù)集能夠裝入 CPU 內(nèi)存時(shí)應(yīng)優(yōu)先選用。兩類數(shù)據(jù)集共享一套統(tǒng)一的目錄約定沿用torchvision的慣例每個(gè)數(shù)據(jù)集接收一個(gè)root文件夾作為存儲(chǔ)根目錄并在其下拆分為兩個(gè)子目錄raw_dir默認(rèn)root/raw存放下載得到的原始數(shù)據(jù)processed_dir默認(rèn)root/processed存放加工后的數(shù)據(jù)。這兩個(gè)路徑分別由 dataset.py 中的raw_dir與processed_dir屬性計(jì)算得出你也可以像 Planetoid 那樣覆寫它們以支持按數(shù)據(jù)集名如root/Cora/raw分目錄存放。此外每個(gè)數(shù)據(jù)集都可以接收三個(gè)默認(rèn)均為None的回調(diào)函數(shù)參數(shù)作用時(shí)機(jī)典型用途transform每次訪問數(shù)據(jù)對(duì)象之前動(dòng)態(tài)執(zhí)行數(shù)據(jù)增強(qiáng)如隨機(jī)擾動(dòng)、掩碼pre_transform數(shù)據(jù)對(duì)象保存到磁盤之前執(zhí)行只需執(zhí)行一次的重量級(jí)預(yù)計(jì)算如添加虛擬節(jié)點(diǎn)、圖歸一化pre_filter保存前手動(dòng)過濾數(shù)據(jù)對(duì)象限制數(shù)據(jù)對(duì)象屬于特定類別等篩選邏輯三者具體的語義在 in_memory_dataset.py 與 dataset.py 的 docstring 中有完整定義。其中pre_transform與pre_filter的執(zhí)行結(jié)果還會(huì)被序列化為processed_dir下的pre_transform.pt與pre_filter.pt文件當(dāng)再次實(shí)例化數(shù)據(jù)集時(shí)如果傳入的pre_transform/pre_filter與磁盤上記錄的不一致PyG 會(huì)發(fā)出警告提示你顯式傳入force_reloadTrue以重新處理見 dataset.py 的_process實(shí)現(xiàn)。這一機(jī)制有效避免了加工過但鉤子函數(shù)已變化的靜默錯(cuò)誤。創(chuàng)建 In-Memory 數(shù)據(jù)集四個(gè)必須實(shí)現(xiàn)的方法要?jiǎng)?chuàng)建一個(gè)InMemoryDataset你需要實(shí)現(xiàn)四個(gè)基礎(chǔ)方法其抽象定義見 in_memory_dataset.pyraw_file_names屬性raw_dir中必須存在的文件列表用于判斷是否可以跳過下載processed_file_names屬性processed_dir中必須存在的文件列表用于判斷是否可以跳過處理download()把原始數(shù)據(jù)下載到raw_dirprocess()讀取原始數(shù)據(jù)、加工并保存到processed_dir。生命周期下載與處理如何被自動(dòng)觸發(fā)Dataset.__init__見 dataset.py在構(gòu)造時(shí)會(huì)依次執(zhí)行if self.has_download: self._download() if self.has_process: self._process()其中has_download/has_process通過overrides_methoddataset.py檢測(cè)子類是否真正覆寫了對(duì)應(yīng)方法_download僅在raw_paths中的文件尚不存在時(shí)才調(diào)用download()_process則在force_reloadFalse且processed_paths文件已存在時(shí)直接跳過處理。這意味著只要raw_file_names/processed_file_names返回的文件都已就位重復(fù)實(shí)例化數(shù)據(jù)集不會(huì)重復(fù)下載或重復(fù)處理——這是 PyG 數(shù)據(jù)集天然具備的緩存能力。完整實(shí)現(xiàn)示例官方教程給出了一個(gè)最小化但完整的InMemoryDataset實(shí)現(xiàn)import torch from torch_geometric.data import InMemoryDataset, download_url class MyOwnDataset(InMemoryDataset): def __init__(self, root, transformNone, pre_transformNone, pre_filterNone): super().__init__(root, transform, pre_transform, pre_filter) self.load(self.processed_paths[0]) # For PyG2.4: # self.data, self.slices torch.load(self.processed_paths[0]) property def raw_file_names(self): return [some_file_1, some_file_2, ...] property def processed_file_names(self): return [data.pt] def download(self): # Download to self.raw_dir. download_url(url, self.raw_dir) ... def process(self): # Read data into huge Data list. data_list [...] if self.pre_filter is not None: data_list [data for data in data_list if self.pre_filter(data)] if self.pre_transform is not None: data_list [self.pre_transform(data) for data in data_list] self.save(data_list, self.processed_paths[0]) # For PyG2.4: # torch.save(self.collate(data_list), self.processed_paths[0])PyG ≥ 2.4 的新機(jī)制save / load 取代手動(dòng) collate torch.load教程中特別注明從 PyG 2.4 起torch.save與InMemoryDataset.collate的功能被統(tǒng)一封裝進(jìn)InMemoryDataset.save而self.data與self.slices也改由InMemoryDataset.load隱式加載。兩者在源碼中的實(shí)現(xiàn)如下in_memory_dataset.pyclassmethod def save(cls, data_list, path): Saves a list of data objects to the file path path. data, slices cls.collate(data_list) fs.torch_save((data.to_dict(), slices, data.__class__), path) def load(self, path, data_clsData): Loads the dataset from the file path path. out fs.torch_load(path) ... if len(out) 2: # Backward compatibility. data, self.slices out else: data, self.slices, data_cls out if not isinstance(data, dict): # Backward compatibility. self.data data else: self.data data_cls.from_dict(data)可以看到save保存的是三元組(data.to_dict(), slices, data.__class__)load在讀取時(shí)會(huì)優(yōu)先兼容舊的二元組格式即 PyG 2.4 的(data, slices)并支持自定義data_cls例如HeteroData的反序列化。因此在自定義數(shù)據(jù)集里__init__結(jié)尾調(diào)用self.load(self.processed_paths[0])即可完成全部加載。collate 與 slices把一個(gè) Python 列表壓縮成一個(gè)對(duì)象教程強(qiáng)調(diào)直接保存一個(gè)巨大的 Python 列表非常緩慢因此我們?cè)诒4媲巴ㄟ^InMemoryDataset.collatein_memory_dataset.py把列表拼接成一個(gè)巨大的Data對(duì)象并額外得到用于還原單個(gè)樣本的slices字典。其底層由 torch_geometric/data/collate.py 的collate函數(shù)完成將所有樣本按屬性如x、edge_index、y縱向拼接為統(tǒng)一表示維護(hù)slice_dict記錄每個(gè)屬性在各樣本間的切片邊界用于從大對(duì)象中重構(gòu)單個(gè)樣本維護(hù)inc_dict記錄各屬性需要累加的量——例如edge_index需要按前序樣本的節(jié)點(diǎn)數(shù)累加偏移見 collate.py 中關(guān)于inc_dict的注釋還原時(shí)再做遞減。InMemoryDataset.get(idx)in_memory_dataset.py正是利用separate與slices從self._data中切出第idx個(gè)樣本并帶有一層緩存_data_list與拷貝保護(hù)當(dāng)數(shù)據(jù)集只有單個(gè)樣本slices is None時(shí)len()返回 1get(0)直接返回_data的淺拷貝。這就是self.data與self.slices兩枚屬性的全部用途——同時(shí)要注意InMemoryDataset.data屬性會(huì)發(fā)出不建議直接訪問內(nèi)部存儲(chǔ)的警告in_memory_dataset.py推薦通過dataset[0]或dataset.x等接口訪問。倉庫中的現(xiàn)成范例KarateClub最簡(jiǎn)單的InMemoryDataset直接在內(nèi)存中構(gòu)造Data(x, edge_index, y, train_mask)最后一行self.data, self.slices self.collate([data])展示了手動(dòng) collate 的經(jīng)典寫法Planetoid展示了raw_dir/processed_dir覆寫、raw_file_names返回帶前綴的文件名列表ind.cora.x等、download中調(diào)用download_url拉取遠(yuǎn)程數(shù)據(jù)以及在__init__里根據(jù)split參數(shù)public/full/random等對(duì)self.data, self.slices self.collate([data])進(jìn)行二次加工FakeDataset用generate_data()批量生成隨機(jī)Data對(duì)象后self.collate(data_list)非常適合在無法聯(lián)網(wǎng)時(shí)快速驗(yàn)證模型與訓(xùn)練流程。其中download_url的實(shí)現(xiàn)見 torch_geometric/data/download.py它會(huì)把 URL 末尾的文件名作為保存名可通過filename參數(shù)覆蓋若文件已存在則直接復(fù)用并打印Using existing file ...下載過程按 10MB 分塊寫入并同時(shí)暴露了download_google_url按 Google Drive 文件 ID 下載。此外 torch_geometric/data/init.py 還提供了extract_tar/extract_zip/extract_bz2/extract_gz等解壓工具download方法里解壓原始?jí)嚎s包時(shí)可以按需調(diào)用。創(chuàng)建大規(guī)模數(shù)據(jù)集Dataset 與按需加載當(dāng)數(shù)據(jù)無法整體裝入內(nèi)存時(shí)應(yīng)使用Dataset基類。它緊密沿襲torchvision數(shù)據(jù)集的概念除了上述四個(gè)方法外還需要額外實(shí)現(xiàn)兩個(gè)方法len()返回?cái)?shù)據(jù)集中樣本的數(shù)量get(idx)實(shí)現(xiàn)加載單個(gè)圖的邏輯。其抽象簽名見 dataset.py。內(nèi)部機(jī)制上Dataset.__getitem__(idx)dataset.py會(huì)先調(diào)用self.get(self.indices()[idx])取回?cái)?shù)據(jù)對(duì)象再按需應(yīng)用transform若傳入切片、列表、torch.Tensor/np.ndarraylong 或 bool 類型等索引則會(huì)走index_select返回?cái)?shù)據(jù)集的子集視圖。因此你只需保證get(idx)足夠高效例如直接torch.load單文件__getitem__、__iter__、len、shuffle、index_select等派生能力便可免費(fèi)獲得。教程給出的Dataset實(shí)現(xiàn)示例如下——每個(gè)圖數(shù)據(jù)對(duì)象在process中單獨(dú)保存為data_{idx}.pt并在get中手動(dòng)加載import os.path as osp import torch from torch_geometric.data import Dataset, download_url class MyOwnDataset(Dataset): def __init__(self, root, transformNone, pre_transformNone, pre_filterNone): super().__init__(root, transform, pre_transform, pre_filter) property def raw_file_names(self): return [some_file_1, some_file_2, ...] property def processed_file_names(self): return [data_1.pt, data_2.pt, ...] def download(self): # Download to self.raw_dir. path download_url(url, self.raw_dir) ... def process(self): idx 0 for raw_path in self.raw_paths: # Read data from raw_path. data Data(...) if self.pre_filter is not None and not self.pre_filter(data): continue if self.pre_transform is not None: data self.pre_transform(data) torch.save(data, osp.join(self.processed_dir, fdata_{idx}.pt)) idx 1 def len(self): return len(self.processed_file_names) def get(self, idx): data torch.load(osp.join(self.processed_dir, fdata_{idx}.pt)) return data注意該例中pre_filter采用不滿足條件就continue跳過保存的方式pre_transform在保存前原地改寫dataprocessed_file_names返回的列表長(zhǎng)度即最終樣本數(shù)因此len()直接取它的長(zhǎng)度即可。此外若你的process會(huì)生成新文件還可以利用raw_paths/processed_paths屬性dataset.py拿到raw_dir/processed_dir下所有文件的絕對(duì)路徑列表與raw_file_names/processed_file_names一一對(duì)應(yīng)。常見問題FAQ如何跳過download和/或process的執(zhí)行只需不覆寫download與process方法即可——has_download/has_process會(huì)返回False構(gòu)造時(shí)便不會(huì)觸發(fā)對(duì)應(yīng)流程。教程給出的寫法是class MyOwnDataset(Dataset): def __init__(self, transformNone, pre_transformNone): super().__init__(None, transform, pre_transform)這里rootNone表示不涉及磁盤讀寫完全在內(nèi)存中工作Dataset.__init__會(huì)把root規(guī)范化為占位符MISSING見 dataset.py。如果只想在已處理數(shù)據(jù)上做隨機(jī)變換甚至不需要繼承數(shù)據(jù)集類。一定要使用這兩套數(shù)據(jù)集接口嗎不必須。正如在原生 PyTorch 中一樣如果你想在飛行中生成合成數(shù)據(jù)、且不需要顯式落盤可以直接把存放Data對(duì)象的普通 Python 列表交給DataLoaderfrom torch_geometric.data import Data from torch_geometric.loader import DataLoader data_list [Data(...), ..., Data(...)] loader DataLoader(data_list, batch_size32)torch_geometric/loader 中的DataLoader會(huì)基于collate自動(dòng)把一批圖拼成Batch對(duì)象這正是快速原型驗(yàn)證的捷徑。另外Dataset還內(nèi)置了get_summary()/print_summary()dataset.py統(tǒng)計(jì)數(shù)據(jù)集概覽、to_datapipe()dataset.py轉(zhuǎn)換為torch.utils.data.DataPipe、以及shuffle(return_perm...)等便捷方法值得在自定義數(shù)據(jù)集中直接繼承復(fù)用。練習(xí)與解答考慮下面這個(gè)由Data對(duì)象列表構(gòu)造的InMemoryDataset對(duì)應(yīng)教程 docs/source/tutorial/create_dataset.rst 的 Exercises 一節(jié)class MyDataset(InMemoryDataset): def __init__(self, root, data_list, transformNone): self.data_list data_list super().__init__(root, transform) self.load(self.processed_paths[0]) property def processed_file_names(self): return data.pt def process(self): self.save(self.data_list, self.processed_paths[0])問題 1self.processed_paths[0]的輸出是什么答案是root目錄下processed子目錄中名為data.pt的絕對(duì)路徑即os.path.join(self.root, processed, data.pt)。依據(jù)是 dataset.py 的processed_paths實(shí)現(xiàn)它把processed_file_names的結(jié)果此處為字符串data.pt會(huì)被to_list包裝成[data.pt]逐一與processed_dir拼接。問題 2InMemoryDataset.save做了什么save是 PyG ≥ 2.4 引入的類方法它首先調(diào)用cls.collate(data_list)把樣本列表拼接為單個(gè)Data對(duì)象并產(chǎn)出slices字典然后把三元組(data.to_dict(), slices, data.__class__)通過fs.torch_save寫入指定路徑in_memory_dataset.py。與之配套的load則讀取該文件并按需恢復(fù)出self.data通過data_cls.from_dict重建對(duì)象與self.slices從而在__init__中完成數(shù)據(jù)集的內(nèi)存化加載。若數(shù)據(jù)集只有一個(gè)樣本collate會(huì)直接返回該樣本、slices為None見 in_memory_dataset.py此時(shí)len()為 1。深入驗(yàn)證測(cè)試與更多擴(kuò)展倉庫的測(cè)試代碼進(jìn)一步印證了上述機(jī)制的預(yù)期行為test/data/test_dataset.py 覆蓋了Dataset的len/get/index_select/shuffle以及切片索引行為test/data/test_inherit.py 驗(yàn)證了raw_file_names等屬性被誤寫為普通方法時(shí)的兼容處理dataset.py 中isinstance(files, Callable)的防御邏輯test/data/test_data.py 與 test/loader/test_dataloader.py 則覆蓋了Data對(duì)象與DataLoader的批量拼接鏈路。當(dāng)內(nèi)存仍顯緊張時(shí)還可以調(diào)用InMemoryDataset.to_on_disk_dataset()in_memory_dataset.py把內(nèi)存數(shù)據(jù)集一鍵轉(zhuǎn)換為基于 SQLite 等后端的 OnDiskDataset適用于分布式訓(xùn)練或共享內(nèi)存受限的硬件環(huán)境——這是教程之外 PyG 為超大數(shù)據(jù)集提供的另一條進(jìn)階路徑。至此從目錄約定、三個(gè)鉤子函數(shù)到InMemoryDataset的四個(gè)方法與save/load/collate底層原理再到Dataset的按需加載與常見問題你已經(jīng)掌握在 PyG 中構(gòu)建自定義圖數(shù)據(jù)集的完整方法論。動(dòng)手實(shí)現(xiàn)你自己的MyOwnDataset配合DataLoader即可無縫接入下游的 GNN 訓(xùn)練與評(píng)測(cè)流程。【免費(fèi)下載鏈接】pytorch_geometricGraph Neural Network Library for PyTorch項(xiàng)目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考