Pytorch Dataset,TensorDataset,Dataloader,Sampler關(guān)系解讀
Dataloader
Dataloader是數(shù)據(jù)加載器,組合數(shù)據(jù)集和采樣器,并在數(shù)據(jù)集上提供單線程或多線程的迭代器。
所以Dataloader的參數(shù)必然需要指定數(shù)據(jù)集Dataset和采樣器Sampler。
class torch.utils.data.DataLoader(dataset, batch_size=1, shuffle=False, sampler=None, num_workers=0, collate_fn=<function default_collate>, pin_memory=False, drop_last=False)
- dataset (Dataset) – 數(shù)據(jù)集。
- batch_size (int, optional) – 每個(gè)batch加載樣本數(shù)。
- shuffle (bool, optional) – True則打亂數(shù)據(jù).
- sampler (Sampler, optional) – 采樣器,如指定則忽略shuffle參數(shù)。
- num_workers (int, optional) – 用多少個(gè)子進(jìn)程加載數(shù)據(jù)。0表示數(shù)據(jù)將在主進(jìn)程中加載
- collate_fn (callable, optional) – 獲取batch數(shù)據(jù)的回調(diào)函數(shù),也就是說(shuō)可以在這個(gè)函數(shù)中修改batch的形式
- pin_memory (bool, optional) –
- drop_last (bool, optional) – 如果數(shù)據(jù)集大小不能被batch size整除,則設(shè)置為True后可刪除最后一個(gè)不完整的batch。如果設(shè)為False并且數(shù)據(jù)集的大小不能被batch size整除,則最后一個(gè)batch將更小。
Dataset和TensorDataset
所有其他數(shù)據(jù)集都應(yīng)該進(jìn)行子類化。所有子類應(yīng)該override __len__ 和 __getitem__ ,前者提供了數(shù)據(jù)集的大小,后者支持整數(shù)索引,范圍從0到len(self)。
TensorDataset是Dataset的子類,已經(jīng)復(fù)寫了 __len__ 和 __getitem__ 方法,只要傳入張量即可,它通過(guò)第一個(gè)維度進(jìn)行索引。

所以TensorDataset說(shuō)白了就是將輸入的tensors捆綁在一起,然后 __len__ 是任何一個(gè)tensor的維度, __getitem__ 表示每個(gè)tensor取相同的索引,然后將這個(gè)結(jié)果組成一個(gè)元組,源碼如下,要好好理解它通過(guò)第一個(gè)維度進(jìn)行索引的意思(針對(duì)tensors里面的每一個(gè)tensor而言)。
class TensorDataset(Dataset): def __init__(self,*tensors): assert all(tensors[0].size(0)==tensor.size(0) for tensor in tensors) self.tensors = tensors def __getitem__(self,index): return tuple(tensor[index] for tensor in self.tensors) def __len__(self): return self.tensors[0].size(0)
Sampler和RandomSampler
Sampler與Dataset類似,是采樣器的基礎(chǔ)類。
每個(gè)采樣器子類必須提供一個(gè) __iter__ 方法,提供一種迭代數(shù)據(jù)集元素的索引的方法,以及返回迭代器長(zhǎng)度的 __len__ 方法。
所以Sampler必然是關(guān)于索引的迭代器,也就是它的輸出是索引。
而RandomSampler與TensorDataset類似,RandomSamper已經(jīng)實(shí)現(xiàn)了 __iter__ 和 __len__ 方法,只需要傳入數(shù)據(jù)集即可。

猜想理解RandomSampler的實(shí)現(xiàn)方式,考慮到這個(gè)類實(shí)現(xiàn)需要傳入Dataset,所以 __len__ 就是Dataset的 __len__ ,然后 __iter__ 就可以隨便搞一個(gè)隨機(jī)函數(shù)對(duì)range(length)隨機(jī)即可。
綜合示例
結(jié)合TensorDataset和RandomSampler使用Dataloader

這里即可理解Dataloader這個(gè)數(shù)據(jù)加載器其實(shí)就是組合數(shù)據(jù)集和采樣器的組合。
所以那就是先根據(jù)Sampler隨機(jī)拿到一個(gè)索引,再用這個(gè)索引到Dataset中取tensors里每個(gè)tensor對(duì)應(yīng)索引的數(shù)據(jù)來(lái)組成一個(gè)元組。
總結(jié)
以上為個(gè)人經(jīng)驗(yàn),希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。
- 使用PyTorch/TensorFlow搭建簡(jiǎn)單全連接神經(jīng)網(wǎng)絡(luò)
- PyTorch使用教程之Tensor包詳解
- 最新tensorflow與pytorch環(huán)境搭建的實(shí)現(xiàn)步驟
- pytorch?tensor合并與分割方式
- Pytorch實(shí)現(xiàn)tensor序列化和并行化的示例詳解
- PyTorch?TensorFlow機(jī)器學(xué)習(xí)框架選擇實(shí)戰(zhàn)
- pytorch中tensorboard安裝及安裝過(guò)程中出現(xiàn)的常見(jiàn)錯(cuò)誤問(wèn)題
- Pytorch之tensorboard無(wú)法啟動(dòng)和顯示問(wèn)題及解決
- PyTorch中tensor[..., 2:4]的實(shí)現(xiàn)示例
相關(guān)文章
基于python實(shí)現(xiàn)PDF分頁(yè)和管理工具開(kāi)發(fā)詳解
本文將詳細(xì)分析一個(gè)使用wxPython開(kāi)發(fā)的PDF分離和管理工具,該工具能夠?qū)DF文件按頁(yè)分離,提供預(yù)覽功能,并支持別名管理系統(tǒng),感興趣的小伙伴可以了解下2025-09-09
Python如何設(shè)置utf-8為默認(rèn)編碼的問(wèn)題
這篇文章主要介紹了Python如何設(shè)置utf-8為默認(rèn)編碼的問(wèn)題,具有很好的參考價(jià)值,希望對(duì)大家有所幫助,如有錯(cuò)誤或未考慮完全的地方,望不吝賜教2024-06-06
python政策網(wǎng)字體反爬實(shí)例(附完整代碼)
大家好,本篇文章主要講的是python政策網(wǎng)字體反爬實(shí)例(附完整代碼),感興趣的同學(xué)趕快來(lái)看一看吧,對(duì)你有幫助的話記得收藏一下2022-01-01
Python實(shí)現(xiàn)Web指紋識(shí)別實(shí)例
這篇文章主要來(lái)帶大家探索Web指紋識(shí)別:了解主流識(shí)別方式,從標(biāo)題到指紋讀取網(wǎng)站信息的簡(jiǎn)單方法,揭秘Web指紋識(shí)別 關(guān)鍵字、哈希和URL的魔力2023-10-10
聊聊Python中的浮點(diǎn)數(shù)運(yùn)算不準(zhǔn)確問(wèn)題
這篇文章主要介紹了聊聊Python中的浮點(diǎn)數(shù)運(yùn)算不準(zhǔn)確問(wèn)題,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧2021-03-03
Django命名URL和反向解析URL實(shí)現(xiàn)解析
這篇文章主要介紹了Django命名URL和反向解析URL實(shí)現(xiàn)解析,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下2019-08-08
Python爬蟲scrapy框架Cookie池(微博Cookie池)的使用
這篇文章主要介紹了Python爬蟲scrapy框架Cookie池(微博Cookie池)的使用,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧2021-01-01

