Python中torch.load()加載模型以及其map_location參數(shù)詳解
參考
torch.load()
函數(shù)格式為:torch.load(f, map_location=None, pickle_module=pickle, **pickle_load_args),一般我們使用的時(shí)候,基本只使用前兩個(gè)參數(shù)。
模型的保存
模型保存有兩種形式,一種是保存模型的state_dict(),只是保存模型的參數(shù)。那么加載時(shí)需要先創(chuàng)建一個(gè)模型的實(shí)例model,之后通過torch.load()將保存的模型參數(shù)加載進(jìn)來,得到dict,再通過model.load_state_dict(dict)將模型的參數(shù)更新。
另一種是將整個(gè)模型保存下來,之后加載的時(shí)候只需要通過torch.load()將模型加載,即可返回一個(gè)加載好的模型。
具體可參考:PyTorch模型的保存與加載。
模型加載中的map_location參數(shù)
具體來說,map_location參數(shù)是用于重定向,比如此前模型的參數(shù)是在cpu中的,我們希望將其加載到cuda:0中。或者我們有多張卡,那么我們就可以將卡1中訓(xùn)練好的模型加載到卡2中,這在數(shù)據(jù)并行的分布式深度學(xué)習(xí)中可能會(huì)用到。
首先定義一個(gè)AlexNet,并使用cuda:0將其訓(xùn)練了一個(gè)貓狗分類,之后把模型存儲(chǔ)起來。
map_location=None
我們先把state_dict加載進(jìn)來。
model_path = "./cuda_model.pth" model = torch.load(model_path) print(next(model.parameters()).device)
結(jié)果為:
cuda:0
因?yàn)楸4娴臅r(shí)候就是模型就是cuda:0的,所以加載進(jìn)來也是。
map_location=torch.device()
model_path = "./cuda_model.pth"
model = torch.load(model_path, map_location=torch.device('cpu'))
print(next(model.parameters()).device)
結(jié)果為:
cpu
模型從cuda:0變成了cpu。
map_location={xx:xx}
model_path = "./cuda_model.pth"
model = torch.load(model_path, map_location={'cuda:0':'cuda:1'})
print(next(model.parameters()).device)
結(jié)果為:
cuda:1
模型從cuda:0變成了cuda:1。
model_path = "./cuda_model.pth"
model = torch.load(model_path, map_location={'cuda:2':'cpu'})
print(next(model.parameters()).device)
結(jié)果為:
cuda:0
模型還是cuda:0,并沒有變成cpu。因?yàn)檫@個(gè)map_location的映射是不對(duì)的,原始的模型就是cuda:0,而映射是cuda:2到cpu,是不對(duì)的。這種情況下,map_location返回None,也就是和不加map_location相同。
總結(jié)
到此這篇關(guān)于Python中torch.load()加載模型以及其map_location參數(shù)詳解的文章就介紹到這了,更多相關(guān)torch.load()加載模型map_location參數(shù)內(nèi)容請(qǐng)搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!
相關(guān)文章
關(guān)于Python中的向量相加和numpy中的向量相加效率對(duì)比
今天小編就為大家分享一篇關(guān)于Python中的向量相加和numpy中的向量相加效率對(duì)比,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過來看看吧2019-08-08
關(guān)于jupyter lab安裝及導(dǎo)入tensorflow找不到模塊的問題
這篇文章主要介紹了關(guān)于jupyter lab安裝及導(dǎo)入tensorflow找不到模塊的問題,本文給大家介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或工作具有一定的參考借鑒價(jià)值,需要的朋友可以參考下2021-03-03
Python光學(xué)仿真學(xué)習(xí)衍射算法初步理解
這篇文章主要為大家介紹了Python光學(xué)仿真學(xué)習(xí)中對(duì)衍射算法的初步理解,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步2021-10-10
Python協(xié)程的實(shí)現(xiàn)方式小結(jié)
協(xié)程是Python中強(qiáng)大的并發(fā)編程工具,允許開發(fā)者編寫異步代碼以提高程序的性能和效率,在本文中,我們將深入探討Python中協(xié)程的實(shí)現(xiàn)方式,包括生成器、asyncio庫和async/await關(guān)鍵字,我們還會(huì)提供詳細(xì)的示例代碼,幫助您理解和應(yīng)用協(xié)程,需要的朋友可以參考下2023-11-11

