pytorch 實(shí)現(xiàn)多個(gè)Dataloader同時(shí)訓(xùn)練
看代碼吧~
如果兩個(gè)dataloader的長(zhǎng)度不一樣,那就加個(gè):
from itertools import cycle
僅使用zip,迭代器將在長(zhǎng)度等于最小數(shù)據(jù)集的長(zhǎng)度時(shí)耗盡。 但是,使用cycle時(shí),我們將再次重復(fù)最小的數(shù)據(jù)集,除非迭代器查看最大數(shù)據(jù)集中的所有樣本。
補(bǔ)充:pytorch技巧:自定義數(shù)據(jù)集 torch.utils.data.DataLoader 及Dataset的使用
本博客中有可直接運(yùn)行的例子,便于直觀的理解,在torch環(huán)境中運(yùn)行即可。
1. 數(shù)據(jù)傳遞機(jī)制
在 pytorch 中數(shù)據(jù)傳遞按一下順序:
1、創(chuàng)建 datasets ,也就是所需要讀取的數(shù)據(jù)集。
2、把 datasets 傳入DataLoader。
3、DataLoader迭代產(chǎn)生訓(xùn)練數(shù)據(jù)提供給模型。
2. torch.utils.data.Dataset
Pytorch提供兩種數(shù)據(jù)集:
Map式數(shù)據(jù)集 Iterable式數(shù)據(jù)集。其中Map式數(shù)據(jù)集繼承torch.utils.data.Dataset,Iterable式數(shù)據(jù)集繼承torch.utils.data.IterableDataset。
本文只介紹 Map式數(shù)據(jù)集。
一個(gè)Map式的數(shù)據(jù)集必須要重寫(xiě) __getitem__(self, index)、 __len__(self) 兩個(gè)方法,用來(lái)表示從索引到樣本的映射(Map)。 __getitem__(self, index)按索引映射到對(duì)應(yīng)的數(shù)據(jù), __len__(self)則會(huì)返回這個(gè)數(shù)據(jù)集的長(zhǎng)度。
基本格式如下:
import torch.utils.data as data class VOCDetection(data.Dataset): ''' 必須繼承data.Dataset類 ''' def __init__(self): ''' 在這里進(jìn)行初始化,一般是初始化文件路徑或文件列表 ''' pass def __getitem__(self, index): ''' 1. 按照index,讀取文件中對(duì)應(yīng)的數(shù)據(jù) (讀取一個(gè)數(shù)據(jù)!?。?!我們常讀取的數(shù)據(jù)是圖片,一般我們送入模型的數(shù)據(jù)成批的,但在這里只是讀取一張圖片,成批后面會(huì)說(shuō)到) 2. 對(duì)讀取到的數(shù)據(jù)進(jìn)行數(shù)據(jù)增強(qiáng) (數(shù)據(jù)增強(qiáng)是深度學(xué)習(xí)中經(jīng)常用到的,可以提高模型的泛化能力) 3. 返回?cái)?shù)據(jù)對(duì) (一般我們要返回 圖片,對(duì)應(yīng)的標(biāo)簽) 在這里因?yàn)槲覜](méi)有寫(xiě)完整的代碼,返回值用 0 代替 ''' return 0 def __len__(self): ''' 返回?cái)?shù)據(jù)集的長(zhǎng)度 ''' return 0
可直接運(yùn)行的例子:
import torch.utils.data as data import numpy as np x = np.array(range(80)).reshape(8, 10) # 模擬輸入, 8個(gè)樣本,每個(gè)樣本長(zhǎng)度為10 y = np.array(range(8)) # 模擬對(duì)應(yīng)樣本的標(biāo)簽, 8個(gè)標(biāo)簽 class Mydataset(data.Dataset): def __init__(self, x, y): self.x = x self.y = y self.idx = list() for item in x: self.idx.append(item) pass def __getitem__(self, index): input_data = self.idx[index] #可繼續(xù)進(jìn)行數(shù)據(jù)增強(qiáng),這里沒(méi)有進(jìn)行數(shù)據(jù)增強(qiáng)操作 target = self.y[index] return input_data, target def __len__(self): return len(self.idx) datasets = Mydataset(x, y) # 初始化 print(datasets.__len__()) # 調(diào)用__len__() 返回?cái)?shù)據(jù)的長(zhǎng)度 for i in range(len(y)): input_data, target = datasets.__getitem__(i) # 調(diào)用__getitem__(index) 返回讀取的數(shù)據(jù)對(duì) print('input_data%d =' % i, input_data) print('target%d = ' % i, target)
結(jié)果如下:
3. torch.utils.data.DataLoader
PyTorch中數(shù)據(jù)讀取的一個(gè)重要接口是 torch.utils.data.DataLoader。
該接口主要用來(lái)將自定義的數(shù)據(jù)讀取接口的輸出或者PyTorch已有的數(shù)據(jù)讀取接口的輸入按照batch_size封裝成Tensor,后續(xù)只需要再包裝成Variable即可作為模型的輸入。
torch.utils.data.DataLoader(onject)的可用參數(shù)如下:
1.dataset(Dataset)
: 數(shù)據(jù)讀取接口,該輸出是torch.utils.data.Dataset類的對(duì)象(或者繼承自該類的自定義類的對(duì)象)。
2.batch_size (int, optional)
: 批訓(xùn)練數(shù)據(jù)量的大小,根據(jù)具體情況設(shè)置即可。一般為2的N次方(默認(rèn):1)
3.shuffle (bool, optional)
:是否打亂數(shù)據(jù),一般在訓(xùn)練數(shù)據(jù)中會(huì)采用。(默認(rèn):False)
4.sampler (Sampler, optional)
:從數(shù)據(jù)集中提取樣本的策略。如果指定,“shuffle”必須為false。我沒(méi)有用過(guò),不太了解。
5.batch_sampler (Sampler, optional)
:和batch_size、shuffle等參數(shù)互斥,一般用默認(rèn)。
6.num_workers
:這個(gè)參數(shù)必須大于等于0,為0時(shí)默認(rèn)使用主線程讀取數(shù)據(jù),其他大于0的數(shù)表示通過(guò)多個(gè)進(jìn)程來(lái)讀取數(shù)據(jù),可以加快數(shù)據(jù)讀取速度,一般設(shè)置為2的N次方,且小于batch_size(默認(rèn):0)
7.collate_fn (callable, optional)
: 合并樣本清單以形成小批量。用來(lái)處理不同情況下的輸入dataset的封裝。
8.pin_memory (bool, optional)
:如果設(shè)置為T(mén)rue,那么data loader將會(huì)在返回它們之前,將tensors拷貝到CUDA中的固定內(nèi)存中.
9.drop_last (bool, optional)
: 如果數(shù)據(jù)集大小不能被批大小整除,則設(shè)置為“true”以除去最后一個(gè)未完成的批。如果“false”那么最后一批將更小。(默認(rèn):false)
10.timeout(numeric, optional)
:設(shè)置數(shù)據(jù)讀取時(shí)間限制,超過(guò)這個(gè)時(shí)間還沒(méi)讀取到數(shù)據(jù)的話就會(huì)報(bào)錯(cuò)。(默認(rèn):0)
11.worker_init_fn (callable, optional)
: 每個(gè)worker初始化函數(shù)(默認(rèn):None)
可直接運(yùn)行的例子:
import torch.utils.data as data import numpy as np x = np.array(range(80)).reshape(8, 10) # 模擬輸入, 8個(gè)樣本,每個(gè)樣本長(zhǎng)度為10 y = np.array(range(8)) # 模擬對(duì)應(yīng)樣本的標(biāo)簽, 8個(gè)標(biāo)簽 class Mydataset(data.Dataset): def __init__(self, x, y): self.x = x self.y = y self.idx = list() for item in x: self.idx.append(item) pass def __getitem__(self, index): input_data = self.idx[index] target = self.y[index] return input_data, target def __len__(self): return len(self.idx) if __name__ ==('__main__'): datasets = Mydataset(x, y) # 初始化 dataloader = data.DataLoader(datasets, batch_size=4, num_workers=2) for i, (input_data, target) in enumerate(dataloader): print('input_data%d' % i, input_data) print('target%d' % i, target)
結(jié)果如下:(注意看類別,DataLoader把數(shù)據(jù)封裝為T(mén)ensor)
以上為個(gè)人經(jīng)驗(yàn),希望能給大家一個(gè)參考,也希望大家多多支持腳本之家。
相關(guān)文章
Pandas merge合并兩個(gè)DataFram的實(shí)現(xiàn)
本文主要介紹了Pandas merge合并兩個(gè)DataFram的實(shí)現(xiàn),文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧2023-03-03python 解決flask uwsgi 獲取不到全局變量的問(wèn)題
今天小編就為大家分享一篇python 解決flask uwsgi 獲取不到全局變量的問(wèn)題,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧2019-12-12pycharm遠(yuǎn)程調(diào)試openstack的圖文教程
這篇文章主要為大家詳細(xì)介紹了pycharm遠(yuǎn)程調(diào)試openstack的圖文教程,具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下2017-11-11Python?OpenCV的基本使用及相關(guān)函數(shù)
這篇文章主要介紹了Python-OpenCV的基本使用和相關(guān)函數(shù)介紹,主要包括圖像的讀取保存圖像展示問(wèn)題,結(jié)合實(shí)例代碼給大家介紹的非常詳細(xì),需要的朋友可以參考下2022-05-05python騰訊語(yǔ)音合成實(shí)現(xiàn)過(guò)程解析
這篇文章主要介紹了python騰訊語(yǔ)音合成實(shí)現(xiàn)過(guò)程解析,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下2019-08-08pytorch 實(shí)現(xiàn)tensor與numpy數(shù)組轉(zhuǎn)換
今天小編就為大家分享一篇使用pytorch 實(shí)現(xiàn)tensor與numpy數(shù)組轉(zhuǎn)換,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧2019-12-12一百行python代碼將圖片轉(zhuǎn)成字符畫(huà)
這篇文章主要為大家詳細(xì)介紹了一百行python代碼將圖片轉(zhuǎn)成字符畫(huà) ,文中示例代碼介紹的非常詳細(xì),具有一定的參考價(jià)值,感興趣的小伙伴們可以參考一下2018-11-11Python中卷積神經(jīng)網(wǎng)絡(luò)(CNN)入門(mén)教程分分享
卷積神經(jīng)網(wǎng)絡(luò)(Convolutional Neural Networks, CNN)是一類特別適用于處理圖像數(shù)據(jù)的深度學(xué)習(xí)模型,本文介紹了如何使用Keras創(chuàng)建一個(gè)簡(jiǎn)單的CNN模型,并用它對(duì)手寫(xiě)數(shù)字進(jìn)行分類,需要的可以參考一下2023-05-05