一文弄懂Pytorch的DataLoader, DataSet, Sampler之間的關(guān)系
以下內(nèi)容都是針對(duì)Pytorch 1.0-1.1介紹。
很多文章都是從Dataset等對(duì)象自下往上進(jìn)行介紹,但是對(duì)于初學(xué)者而言,其實(shí)這并不好理解,因?yàn)橛械臅r(shí)候會(huì)不自覺地陷入到一些細(xì)枝末節(jié)中去,而不能把握重點(diǎn),所以本文將會(huì)自上而下地對(duì)Pytorch數(shù)據(jù)讀取方法進(jìn)行介紹。
自上而下理解三者關(guān)系
首先我們看一下DataLoader.next的源代碼長(zhǎng)什么樣,為方便理解我只選取了num_works為0的情況(num_works簡(jiǎn)單理解就是能夠并行化地讀取數(shù)據(jù))。
class DataLoader(object): ... def __next__(self): if self.num_workers == 0: indices = next(self.sample_iter) # Sampler batch = self.collate_fn([self.dataset[i] for i in indices]) # Dataset if self.pin_memory: batch = _utils.pin_memory.pin_memory_batch(batch) return batch
在閱讀上面代碼前,我們可以假設(shè)我們的數(shù)據(jù)是一組圖像,每一張圖像對(duì)應(yīng)一個(gè)index,那么如果我們要讀取數(shù)據(jù)就只需要對(duì)應(yīng)的index即可,即上面代碼中的indices
,而選取index的方式有多種,有按順序的,也有亂序的,所以這個(gè)工作需要Sampler
完成,現(xiàn)在你不需要具體的細(xì)節(jié),后面會(huì)介紹,你只需要知道DataLoader和Sampler在這里產(chǎn)生關(guān)系。
那么Dataset和DataLoader在什么時(shí)候產(chǎn)生關(guān)系呢?沒錯(cuò)就是下面一行。我們已經(jīng)拿到了indices,那么下一步我們只需要根據(jù)index對(duì)數(shù)據(jù)進(jìn)行讀取即可了。
再下面的if
語(yǔ)句的作用簡(jiǎn)單理解就是,如果pin_memory=True
,那么Pytorch會(huì)采取一系列操作把數(shù)據(jù)拷貝到GPU,總之就是為了加速。
綜上可以知道DataLoader,Sampler和Dataset三者關(guān)系如下:
在閱讀后文的過程中,你始終需要將上面的關(guān)系記在心里,這樣能幫助你更好地理解。
Sampler
參數(shù)傳遞
要更加細(xì)致地理解Sampler原理,我們需要先閱讀一下DataLoader 的源代碼,如下:
class DataLoader(object): def __init__(self, dataset, batch_size=1, shuffle=False, sampler=None, batch_sampler=None, num_workers=0, collate_fn=default_collate, pin_memory=False, drop_last=False, timeout=0, worker_init_fn=None)
可以看到初始化參數(shù)里有兩種sampler:sampler
和batch_sampler
,都默認(rèn)為None
。前者的作用是生成一系列的index,而batch_sampler則是將sampler生成的indices打包分組,得到一個(gè)又一個(gè)batch的index。例如下面示例中,BatchSampler
將SequentialSampler
生成的index按照指定的batch size分組。
>>>in : list(BatchSampler(SequentialSampler(range(10)), batch_size=3, drop_last=False)) >>>out: [[0, 1, 2], [3, 4, 5], [6, 7, 8], [9]]
Pytorch中已經(jīng)實(shí)現(xiàn)的Sampler
有如下幾種:
SequentialSampler
RandomSampler
WeightedSampler
SubsetRandomSampler
需要注意的是DataLoader的部分初始化參數(shù)之間存在互斥關(guān)系,這個(gè)你可以通過閱讀源碼更深地理解,這里只做總結(jié):
- 如果你自定義了batch_sampler,那么這些參數(shù)都必須使用默認(rèn)值:batch_size, shuffle,sampler,drop_last.
- 如果你自定義了sampler,那么shuffle需要設(shè)置為False
- 如果sampler和batch_sampler都為None,那么batch_sampler使用Pytorch已經(jīng)實(shí)現(xiàn)好的BatchSampler,而sampler分兩種情況:
- 若shuffle=True,則sampler=RandomSampler(dataset)
- 若shuffle=False,則sampler=SequentialSampler(dataset)
如何自定義Sampler和BatchSampler?
仔細(xì)查看源代碼其實(shí)可以發(fā)現(xiàn),所有采樣器其實(shí)都繼承自同一個(gè)父類,即Sampler
,其代碼定義如下:
class Sampler(object): r"""Base class for all Samplers. Every Sampler subclass has to provide an :meth:`__iter__` method, providing a way to iterate over indices of dataset elements, and a :meth:`__len__` method that returns the length of the returned iterators. .. note:: The :meth:`__len__` method isn't strictly required by :class:`~torch.utils.data.DataLoader`, but is expected in any calculation involving the length of a :class:`~torch.utils.data.DataLoader`. """ def __init__(self, data_source): pass def __iter__(self): raise NotImplementedError def __len__(self): return len(self.data_source)
所以你要做的就是定義好__iter__(self)
函數(shù),不過要注意的是該函數(shù)的返回值需要是可迭代的。例如SequentialSampler
返回的是iter(range(len(self.data_source)))
。
另外BatchSampler
與其他Sampler的主要區(qū)別是它需要將Sampler作為參數(shù)進(jìn)行打包,進(jìn)而每次迭代返回以batch size為大小的index列表。也就是說在后面的讀取數(shù)據(jù)過程中使用的都是batch sampler。
Dataset
Dataset定義方式如下:
class Dataset(object): def __init__(self): ... def __getitem__(self, index): return ... def __len__(self): return ...
上面三個(gè)方法是最基本的,其中__getitem__
是最主要的方法,它規(guī)定了如何讀取數(shù)據(jù)。但是它又不同于一般的方法,因?yàn)樗莗ython built-in方法,其主要作用是能讓該類可以像list一樣通過索引值對(duì)數(shù)據(jù)進(jìn)行訪問。假如你定義好了一個(gè)dataset,那么你可以直接通過dataset[0]
來訪問第一個(gè)數(shù)據(jù)。在此之前我一直沒弄清楚__getitem__
是什么作用,所以一直不知道該怎么進(jìn)入到這個(gè)函數(shù)進(jìn)行調(diào)試?,F(xiàn)在如果你想對(duì)__getitem__
方法進(jìn)行調(diào)試,你可以寫一個(gè)for循環(huán)遍歷dataset來進(jìn)行調(diào)試了,而不用構(gòu)建dataloader等一大堆東西了,建議學(xué)會(huì)使用ipdb
這個(gè)庫(kù),非常實(shí)用?。?!以后有時(shí)間再寫一篇ipdb的使用教程。另外,其實(shí)我們通過最前面的Dataloader的__next__
函數(shù)可以看到DataLoader對(duì)數(shù)據(jù)的讀取其實(shí)就是用了for循環(huán)來遍歷數(shù)據(jù),不用往上翻了,我直接復(fù)制了一遍,如下:
class DataLoader(object): ... def __next__(self): if self.num_workers == 0: indices = next(self.sample_iter) batch = self.collate_fn([self.dataset[i] for i in indices]) # this line if self.pin_memory: batch = _utils.pin_memory.pin_memory_batch(batch) return batch
我們仔細(xì)看可以發(fā)現(xiàn),前面還有一個(gè)self.collate_fn
方法,這個(gè)是干嘛用的呢?在介紹前我們需要知道每個(gè)參數(shù)的意義:
indices
: 表示每一個(gè)iteration,sampler返回的indices,即一個(gè)batch size大小的索引列表self.dataset[i]
: 前面已經(jīng)介紹了,這里就是對(duì)第i個(gè)數(shù)據(jù)進(jìn)行讀取操作,一般來說self.dataset[i]=(img, label)
看到這不難猜出collate_fn
的作用就是將一個(gè)batch的數(shù)據(jù)進(jìn)行合并操作。默認(rèn)的collate_fn
是將img和label分別合并成imgs和labels,所以如果你的__getitem__
方法只是返回 img, label
,那么你可以使用默認(rèn)的collate_fn
方法,但是如果你每次讀取的數(shù)據(jù)有img, box, label
等等,那么你就需要自定義collate_fn
來將對(duì)應(yīng)的數(shù)據(jù)合并成一個(gè)batch數(shù)據(jù),這樣方便后續(xù)的訓(xùn)練步驟。
到此這篇關(guān)于一文弄懂Pytorch的DataLoader, DataSet, Sampler之間的關(guān)系的文章就介紹到這了,更多相關(guān)Pytorch DataLoader DataSet Sampler內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
相關(guān)文章
Python Word文件自動(dòng)化實(shí)戰(zhàn)之簡(jiǎn)歷篩選
本文將利用Python自動(dòng)化做一個(gè)具有實(shí)操性的小練習(xí),即通過讀取簡(jiǎn)歷來篩選出符合招聘條件的簡(jiǎn)歷。文中的示例代碼講解詳細(xì),感興趣的小伙伴可以跟隨小編一起學(xué)習(xí)一下2022-05-05Python 轉(zhuǎn)換文本編碼實(shí)現(xiàn)解析
這篇文章主要介紹了Python 轉(zhuǎn)換文本編碼實(shí)現(xiàn)解析,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值2019-08-08Python 繪圖庫(kù) Matplotlib 入門教程
Matplotlib是一個(gè)Python語(yǔ)言的2D繪圖庫(kù),它支持各種平臺(tái),并且功能強(qiáng)大,能夠輕易繪制出各種專業(yè)的圖像。本文是對(duì)Python 繪圖庫(kù) Matplotlib 入門教程,感興趣的朋友跟隨腳本之家小編一起學(xué)習(xí)吧2018-04-04Django基于Models定制Admin后臺(tái)實(shí)現(xiàn)過程解析
這篇文章主要介紹了Django基于Models定制Admin后臺(tái)實(shí)現(xiàn)過程解析,文中通過示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下2020-11-11通過Python編程將CSV文件導(dǎo)出為PDF文件的方法
CSV文件通常用于存儲(chǔ)大量的數(shù)據(jù),而PDF文件則是一種通用的文檔格式,便于與他人共享和打印,將CSV文件轉(zhuǎn)換成PDF文件可以幫助我們更好地管理和展示數(shù)據(jù),本文將介紹如何通過Python編程將CSV文件導(dǎo)出為PDF文件,需要的朋友可以參考下2024-06-06