<em>Mac</em>Book项目 2009年学校开始实施<em>Mac</em>Book项目,所有师生配备一本<em>Mac</em>Book,并同步更新了校园无线网络。学校每周进行电脑技术更新,每月发送技术支持资料,极大改变了教学及学习方式。因此2011
2021-06-01 09:32:01
分為幾個步驟
資料準備以及載入資料庫–>資料載入器的呼叫或者設計–>批次呼叫進行訓練或者其他作用
直接讀取了x和y的資料變數,對比後面的就從把對應的路徑寫進了文字檔案中,通過載入器進行讀取
x = torch.linspace(1, 10, 10) # 訓練資料 linspace返回一個一維的張量,(最小值,最大值,多少個數) print(x) y = torch.linspace(10, 1, 10) # 標籤 print(y)
將資料載入進資料庫
輸出的結果是<torch.utils.data.dataset.TensorDataset object at 0x00000145BD93F1C0>
,需要使用載入器進行載入,才能迭代遍歷
import torch.utils.data as Data torch_dataset = Data.TensorDataset(x, y) # 對給定的 tensor 資料,將他們包裝成 dataset #輸出的結果是<torch.utils.data.dataset.TensorDataset object at 0x00000145BD93F1C0>,需要使用載入器進行載入,才能迭代遍歷 print(torch_dataset)
所以要想看裡面的內容,就需要用迭代進行操作或者檢視。
BATCH_SIZE=5 loader = Data.DataLoader(#使用支援的預設的資料集載入的方式 # 從資料庫中每次抽出batch size個樣本 dataset=torch_dataset, # torch TensorDataset format 載入資料集 batch_size=BATCH_SIZE, # mini batch size 5 shuffle=False, # 要不要打亂資料 (打亂比較好) num_workers=2, # 多執行緒來讀資料 ) def show_batch(): for epoch in range(3): for step, (batch_x, batch_y) in enumerate(loader): #載入資料集的時候起的作用很奇怪 # training print("steop:{}, batch_x:{}, batch_y:{}".format(step, batch_x, batch_y)) print("*"*100) if __name__ == '__main__': show_batch()
實現自己的資料集就需要完成對dataset類的過載。這個類的過載完成幾個函數的作用
__init__()
__getitem__
__len__
基本的資料集的方法就是完成以上步驟,但是可以想想資料集通常是一些圖片和標籤組成,而這些資料集以及標籤是儲存在計算機上,具有相對應的位置,那麼直接存取對應的位置因為是在資料夾下需要進行遍歷等一系列操作,而且這就顯得和dataset類沒有解耦,因為有時候在這些位置的操作可能會有一些特殊操作,所以如果能夠將其位置儲存在文字檔案中可能就會方便很多,所以就採取儲存文字檔案的方式。
# 自定義資料集類 class MyDataset(torch.utils.data.Dataset): def __init__(self, *args): super().__init__() # 初始化資料集包含的資料和標籤 pass def __getitem__(self, index): # 根據索引index從檔案中讀取一個資料 # 對資料預處理 # 返回資料和對應標籤 pass def __len__(self): # 返回資料集的大小 return len()
所以這裡新建一個資料庫就是新建了兩個文字檔案,然後載入器通過文字檔案就將圖片以及label載入進去了。而標準的資料集操作是使用了自帶的資料集介面,在載入的時候也不用再去實現相關的__getitem__方法
以下程式碼用於生成對應的train.txt val.txt
''' 生成訓練集和測試集,儲存在txt檔案中 ''' import os import random train_ratio = 0.6 test_ratio = 1-train_ratio rootdata = r"dataset" #陣列定義 train_list, test_list = [],[] data_list = [] class_flag = -1 # 將絕對路徑載入進陣列中 for a,b,c in os.walk(rootdata):#os.walk可以獲得根路徑、資料夾以及檔案,並會一直進行迭代遍歷下去,直至只有檔案才會結束 print(a) for i in range(len(c)): data_list.append(os.path.join(a,c[i])) for i in range(0,int(len(c)*train_ratio)): train_data = os.path.join(a, c[i])+'t'+str(class_flag)+'n' #class_flag表示分類的類別 train_list.append(train_data) for i in range(int(len(c) * train_ratio),len(c)): test_data = os.path.join(a, c[i]) + 't' + str(class_flag)+'n' test_list.append(test_data) class_flag += 1 print(train_list) # 將陣列的內容打亂順序 random.shuffle(train_list) random.shuffle(test_list) #分別將絕對路徑對應的陣列內容寫進文字檔案裡 with open('train.txt','w',encoding='UTF-8') as f: for train_img in train_list: f.write(str(train_img)) with open('test.txt','w',encoding='UTF-8') as f: for test_img in test_list: f.write(test_img)
初始化資料集中的資料以及標籤、相關變數__init__()
def __init__(self, txt_path, train_flag=True): #初始化圖片對應的變數imgs_info以及一些相關變數 self.imgs_info = self.get_images(txt_path) #imgs_info儲存了圖片以及標籤 self.train_flag = train_flag self.train_tf = transforms.Compose([#對訓練集的圖片進行預處理 transforms.Resize(224), transforms.RandomHorizontalFlip(), transforms.RandomVerticalFlip(), transforms.ToTensor(), transform_BZ ]) self.val_tf = transforms.Compose([#對測試集的圖片進行預處理 transforms.Resize(224), transforms.ToTensor(), transform_BZ ])
返回資料和對應標籤__getitem__
def __getitem__(self, index): img_path, label = self.imgs_info[index] #開啟圖片,並將RGBA轉換為RGB,這裡是通過PIL庫開啟圖片的 img = Image.open(img_path) img = img.convert('RGB') img = self.padding_black(img) #將圖片新增上黑邊的 if self.train_flag: #選擇是訓練集還是測試集 img = self.train_tf(img) else: img = self.val_tf(img) label = int(label) return img, label
返回資料集的大小__len__
def __len__(self): return len(self.imgs_info)
由於前面已經對整合dataset的類進行了實現三種方法,那麼就可以在載入器中進行載入,將載入後的資料傳入到train函數或者test函數都可以
train_dataloader = DataLoader(dataset=train_data, num_workers=4, pin_memory=True, batch_size=batch_size, shuffle=True)
:使用載入器載入資料train(train_dataloader, model, loss_fn, optimizer) test(test_dataloader, model)
:將資料傳入train或者test中進行訓練或者測試if __name__=='__main__': batch_size = 16 # # 給訓練集和測試集分別建立一個資料集載入器 train_data = LoadData("train.txt", True) valid_data = LoadData("test.txt", False) train_dataloader = DataLoader(dataset=train_data, num_workers=4, pin_memory=True, batch_size=batch_size, shuffle=True) test_dataloader = DataLoader(dataset=valid_data, num_workers=4, pin_memory=True, batch_size=batch_size) for X, y in test_dataloader: print("Shape of X [N, C, H, W]: ", X.shape) print("Shape of y: ", y.shape, y.dtype) break
連結: https://pan.baidu.com/s/19Oo87gbcm9e8zvYGkBi95A 提取碼: 2tss
到此這篇關於pytorch載入自己的資料集原始碼分享的文章就介紹到這了,更多相關pytorch載入自己的資料集內容請搜尋it145.com以前的文章或繼續瀏覽下面的相關文章希望大家以後多多支援it145.com!
相關文章
<em>Mac</em>Book项目 2009年学校开始实施<em>Mac</em>Book项目,所有师生配备一本<em>Mac</em>Book,并同步更新了校园无线网络。学校每周进行电脑技术更新,每月发送技术支持资料,极大改变了教学及学习方式。因此2011
2021-06-01 09:32:01
综合看Anker超能充系列的性价比很高,并且与不仅和iPhone12/苹果<em>Mac</em>Book很配,而且适合多设备充电需求的日常使用或差旅场景,不管是安卓还是Switch同样也能用得上它,希望这次分享能给准备购入充电器的小伙伴们有所
2021-06-01 09:31:42
除了L4WUDU与吴亦凡已经多次共事,成为了明面上的厂牌成员,吴亦凡还曾带领20XXCLUB全队参加2020年的一场音乐节,这也是20XXCLUB首次全员合照,王嗣尧Turbo、陈彦希Regi、<em>Mac</em> Ova Seas、林渝植等人全部出场。然而让
2021-06-01 09:31:34
目前应用IPFS的机构:1 谷歌<em>浏览器</em>支持IPFS分布式协议 2 万维网 (历史档案博物馆)数据库 3 火狐<em>浏览器</em>支持 IPFS分布式协议 4 EOS 等数字货币数据存储 5 美国国会图书馆,历史资料永久保存在 IPFS 6 加
2021-06-01 09:31:24
开拓者的车机是兼容苹果和<em>安卓</em>,虽然我不怎么用,但确实兼顾了我家人的很多需求:副驾的门板还配有解锁开关,有的时候老婆开车,下车的时候偶尔会忘记解锁,我在副驾驶可以自己开门:第二排设计很好,不仅配置了一个很大的
2021-06-01 09:30:48
不仅是<em>安卓</em>手机,苹果手机的降价力度也是前所未有了,iPhone12也“跳水价”了,发布价是6799元,如今已经跌至5308元,降价幅度超过1400元,最新定价确认了。iPhone12是苹果首款5G手机,同时也是全球首款5nm芯片的智能机,它
2021-06-01 09:30:45