最新国产好看的视频,伊人天堂AV在线,国产Aaaaaa视频,蜜臀视频在线观看一区,人妻av色图,密臀久久久精品影片,青青视频免费观看毛片,久草在线观看视,国产三级精品色情在线

Pytorch模型的保存/復(fù)用/遷移實現(xiàn)代碼

 更新時間:2023年05月05日 10:32:07   作者:信海  
本文整理了Pytorch框架下模型的保存、復(fù)用、推理、再訓練和遷移等實現(xiàn),本文通過實例代碼給大家介紹的非常詳細,對大家的學習或工作具有一定的參考借鑒價值,需要的朋友可以參考下

本文整理了Pytorch框架下模型的保存、復(fù)用、推理、再訓練和遷移等實現(xiàn)。

模型的保存與復(fù)用

模型定義和參數(shù)打印

# 定義模型結(jié)構(gòu)
class LenNet(nn.Module):
    def __init__(self):
        super(LenNet, self).__init__()
        self.conv = nn.Sequential(  # [batch, 1, 28, 28]
            nn.Conv2d(1, 8, 5, 2),  # [batch, 1, 28, 28]
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),  # [batch, 8, 14, 14]
            nn.Conv2d(8, 16, 5),  # [batch, 16, 10, 10]
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),  # [batch, 16, 5, 5]
        )
        self.fc = nn.Sequential(
            nn.Flatten(),
            nn.Linear(16*5*5, 128),
            nn.ReLU(inplace=True),
            nn.Linear(128, 64),
            nn.ReLU(inplace=True),
            nn.Linear(64, 10)
        )
    def forward(self, X):
        return self.fc(self.conv(X))
# 查看模型參數(shù)
# 網(wǎng)絡(luò)模型中的參數(shù)model.state_dict()是以字典形式保存(實質(zhì)上是collections模塊中的OrderedDict)
model = LenNet()
print("Model's state_dict:")
for param_tensor in model.state_dict():
    print(param_tensor, "\t", model.state_dict()[param_tensor].size())
# 參數(shù)名中的fc和conv前綴是根據(jù)定義nn.Sequential()時的名字所確定。
# 參數(shù)名中的數(shù)字表示每個Sequential()中網(wǎng)絡(luò)層所在的位置。
print(model.state_dict().keys())  # 打印鍵
print(model.state_dict().values())  # 打印值
# 優(yōu)化器optimizer的參數(shù)打印類似
optimizer = torch.optim.SGD(model.parameters(), lr=0.001, momentum=0.9)
print("Optimizer's state_dict:")
for var_name in optimizer.state_dict():
   print(var_name, "\t", optimizer.state_dict()[var_name])

模型保存

import os
# 指定保存的模型名稱時Pytorch官方建議的后綴為.pt或者.pth
model_save_dir = './model_logs/'
model_save_path = os.path.join(model_save_dir, 'LeNet.pt')
torch.save(model.state_dict(), model_save_path)
# 在訓練過程中保存某個條件下的最優(yōu)模型,可以如下操作
best_model_state = deepcopy(model.state_dict()) 
torch.save(best_model_state, model_save_path)
# 下面這種方法是錯誤的,因為best_model_state只是model.state_dict()的引用,會隨著訓練的改變而改變
best_model_state = model.state_dict() 
torch.save(best_model_state, model_save_path)

模型推理

def inference(data_iter, device, model_save_dir):
	model = LeNet()  # 初始化現(xiàn)有模型的權(quán)重參數(shù)
    model.to(device)
    model_save_path = os.path.join(model_save_dir, 'LeNet.pt')
    # 如果本地存在模型,則加載本地模型參數(shù)覆蓋原有模型
    if os.path.exists(model_save_path): 
        loaded_paras = torch.load(model_save_path)
        model.load_state_dict(loaded_paras)
        model.eval()
    with torch.no_grad():  # 開始推理
        acc_sum, n = 0., 0
        for x, y in data_iter:
            x, y = x.to(device), y.to(device)
            logits = model(x)
            acc_sum += (logits.argmax(1) == y).float().sum().item()
            n += len(y)
        print("Accuracy in test data is : ", acc_sum / n)

模型再訓練

class MyModel:
    def __init__(self,
                 batch_size=64,
                 epochs=5,
                 learning_rate=0.001,
                 model_save_dir='./MODEL'):
        self.batch_size = batch_size
        self.epochs = epochs
        self.learning_rate = learning_rate
        self.model_save_dir = model_save_dir
        self.model = LeNet()
    def train(self):
        train_iter, test_iter = load_dataset(self.batch_size)
        # 在訓練過程中只保存網(wǎng)絡(luò)權(quán)重,在再訓練時只載入網(wǎng)絡(luò)權(quán)重參數(shù)初始化網(wǎng)絡(luò)訓練。這里是核心部分,開始。
        if not os.path.exists(self.model_save_dir):
            os.makedirs(self.model_save_dir)
        model_save_path = os.path.join(self.model_save_dir, 'model.pt')
        if os.path.exists(model_save_path):
            loaded_paras = torch.load(model_save_path)
            self.model.load_state_dict(loaded_paras)
            print("#### 成功載入已有模型,進行再訓練...")
        # 結(jié)束  
        optimizer = torch.optim.Adam(self.model.parameters(), lr=self.learning_rate)  
        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        self.model.to(device)
        for epoch in range(self.epochs):
            for i, (x, y) in enumerate(train_iter):
                x, y = x.to(device), y.to(device)
                loss, logits = self.model(x)
                optimizer.zero_grad()
                loss.backward()
                optimizer.step()  
                if i % 100 == 0:
                    acc = (logits.argmax(1) == y).float().mean()
                    print("Epochs[{}/{}]---batch[{}/{}]---acc {:.4}---loss {:.4}".format(
                        epoch, self.epochs, len(train_iter), i, acc, loss.item()))
            print("Epochs[{}/{}]--acc on test {:.4}".format(epoch, self.epochs,
                                                            self.evaluate(test_iter, self.model, device)))
            torch.save(self.model.state_dict(), model_save_path)
    @staticmethod
    def evaluate(data_iter, model, device):
        with torch.no_grad():
            acc_sum, n = 0.0, 0
            for x, y in data_iter:
                x, y = x.to(device), y.to(device)
                logits = model(x)
                acc_sum += (logits.argmax(1) == y).float().sum().item()
                n += len(y)
            return acc_sum / n
# 在保存參數(shù)的時候,將優(yōu)化器參數(shù)、損失值等可一同保存,然后在恢復(fù)模型時連同其它參數(shù)一起恢復(fù)
model_save_path = os.path.join(model_save_dir, 'LeNet.pt')
torch.save({
            'epoch': epoch,
            'model_state_dict': model.state_dict(),
            'optimizer_state_dict': optimizer.state_dict(),
            'loss': loss,
            ...
            }, model_save_path)
# 加載方式如下
checkpoint = torch.load(model_save_path)
model.load_state_dict(checkpoint['model_state_dict'])
optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
epoch = checkpoint['epoch']
loss = checkpoint['loss']

模型遷移

# 定義新模型NewLeNet 和LeNet區(qū)別在于新增了一個全連接層
class NewLenNet(nn.Module):
    def __init__(self):
        super(NewLenNet, self).__init__()
        self.conv = nn.Sequential(  # [batch, 1, 28, 28]
            nn.Conv2d(1, 8, 5, 2),  # [batch, 1, 28, 28]
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),  # [batch, 8, 14, 14]
            nn.Conv2d(8, 16, 5),  # [batch, 16, 10, 10]
            nn.ReLU(inplace=True),
            nn.MaxPool2d(2, 2),  # [batch, 16, 5, 5]
        )
        self.fc = nn.Sequential(
            nn.Flatten(),
            nn.Linear(16*5*5, 128),
            nn.ReLU(inplace=True),
            nn.Linear(128, 64), # 這層以前和LeNet結(jié)構(gòu)一致 可以用LeNet的參數(shù)來進行替換
            nn.ReLU(inplace=True),
            nn.Linear(64, 32),
            nn.ReLU(inplace=True),
            nn.Linear(32, 10)
        )
    def forward(self, X):
        return self.fc(self.conv(X))
# 定義替換函數(shù) 匹配兩個網(wǎng)絡(luò) size相同處地方進行參數(shù)替換
def para_state_dict(model, model_save_dir):
    state_dict = deepcopy(model.state_dict())
    model_save_path = os.path.join(model_save_dir, 'model.pt')
    if os.path.exists(model_save_path):
        loaded_paras = torch.load(model_save_path)
        for key in state_dict:  # 在新的網(wǎng)絡(luò)模型中遍歷對應(yīng)參數(shù)
            if key in loaded_paras and state_dict[key].size() == loaded_paras[key].size():
                print("成功初始化參數(shù):", key)
                state_dict[key] = loaded_paras[key]
    return state_dict
# 更新一下模型遷移后的訓練代碼
def train(self):
        train_iter, test_iter = load_dataset(self.batch_size)
        if not os.path.exists(self.model_save_dir):
            os.makedirs(self.model_save_dir)
        model_save_path = os.path.join(self.model_save_dir, 'model_new.pt')
        old_model = os.path.join(self.model_save_dir, 'LeNet.pt')
        if os.path.exists(old_model):
            state_dict = para_state_dict(self.model, self.model_save_dir)  # 調(diào)用遷移代碼 將LeNet的前幾層參數(shù)遷移到NewLeNet
            self.model.load_state_dict(state_dict)
            print("#### 成功載入已有模型,進行再訓練...")
        optimizer = torch.optim.Adam(self.model.parameters(), lr=self.learning_rate)  
        device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
        self.model.to(device)
        for epoch in range(self.epochs):
            for i, (x, y) in enumerate(train_iter):
                x, y = x.to(device), y.to(device)
                loss, logits = self.model(x)
                optimizer.zero_grad()
                loss.backward()
                optimizer.step()  
                if i % 100 == 0:
                    acc = (logits.argmax(1) == y).float().mean()
                    print("Epochs[{}/{}]---batch[{}/{}]---acc {:.4}---loss {:.4}".format(
                        epoch, self.epochs, len(train_iter), i, acc, loss.item()))
            print("Epochs[{}/{}]--acc on test {:.4}".format(epoch, self.epochs,
                                                            self.evaluate(test_iter, self.model, device)))
            torch.save(self.model.state_dict(), model_save_path)
# 這里更新未進行訓練的推理
def inference(data_iter, device, model_save_dir='./MODEL'):
    model = NewLeNet()  # 初始化現(xiàn)有模型的權(quán)重參數(shù)
    print("初始化參數(shù) conv.0.bias 為:", model.state_dict()['conv.0.bias'])
    model.to(device)
    state_dict = para_state_dict(model, model_save_dir) # 遷移模型參數(shù)
    model.load_state_dict(state_dict)
    model.eval()
    print("載入本地模型重新初始化 conv.0.bias 為:", model.state_dict()['conv.0.bias'])
    with torch.no_grad():
        acc_sum, n = 0.0, 0
        for x, y in data_iter:
            x, y = x.to(device), y.to(device)
            logits = model(x)
            acc_sum += (logits.argmax(1) == y).float().sum().item()
            n += len(y)
        print("Accuracy in test data is :", acc_sum / n)

參考文獻

[1] https://github.com/moon-hotel/DeepLearningWithMe

到此這篇關(guān)于Pytorch模型的保存/復(fù)用/遷移的文章就介紹到這了,更多相關(guān)Pytorch模型保存遷移內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • pycharm使用matplotlib.pyplot不顯示圖形的解決方法

    pycharm使用matplotlib.pyplot不顯示圖形的解決方法

    今天小編就為大家分享一篇pycharm使用matplotlib.pyplot不顯示圖形的解決方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-10-10
  • Python正則表達式匹配ip地址實例

    Python正則表達式匹配ip地址實例

    這篇文章主要介紹了Python正則表達式匹配ip地址實例,通過簡單的實例講述了re模塊的用法,該實例非常具有實用價值,需要的朋友可以參考下
    2014-10-10
  • python讀取excel文件的方法

    python讀取excel文件的方法

    文章介紹了在Python中讀取Excel文件的兩種方法:使用pandas庫和使用openpyxl庫,pandas適合數(shù)據(jù)分析和處理,而openpyxl提供了更多的Excel文件操作功能,感興趣的朋友跟隨小編一起看看吧
    2024-11-11
  • Python 格式化打印json數(shù)據(jù)方法(展開狀態(tài))

    Python 格式化打印json數(shù)據(jù)方法(展開狀態(tài))

    今天小編就為大家分享一篇Python 格式化打印json數(shù)據(jù)方法(展開狀態(tài)),具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2020-02-02
  • Python中語音轉(zhuǎn)文字相關(guān)庫介紹(最新推薦)

    Python中語音轉(zhuǎn)文字相關(guān)庫介紹(最新推薦)

    Python的speech_recognition庫是一個用于語音識別的Python包,它可以使Python程序能夠識別和翻譯來自麥克風、音頻文件或網(wǎng)絡(luò)流的語音,這篇文章主要介紹了Python中語音轉(zhuǎn)文字相關(guān)庫介紹,需要的朋友可以參考下
    2023-05-05
  • celery異步定時任務(wù)訂單定時回滾

    celery異步定時任務(wù)訂單定時回滾

    這篇文章主要為大家介紹了celery異步定時任務(wù)訂單定時回滾的實現(xiàn)示例,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進步早日升職加薪
    2022-04-04
  • Python中用try-except-finally處理異常問題

    Python中用try-except-finally處理異常問題

    這篇文章主要介紹了Python中用try-except-finally處理異常問題,具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-12-12
  • python使用參數(shù)對嵌套字典進行取值的方法

    python使用參數(shù)對嵌套字典進行取值的方法

    這篇文章主要介紹了python使用參數(shù)對嵌套字典進行取值,小編覺得挺不錯的,現(xiàn)在分享給大家,也給大家做個參考。一起跟隨小編過來看看吧
    2019-04-04
  • Python利用pandas處理CSV文件的用法示例

    Python利用pandas處理CSV文件的用法示例

    pandas是一個第三方數(shù)據(jù)分析庫,其集成了大量的數(shù)據(jù)分析工具,可以方便的處理和分析各類數(shù)據(jù),本文將給大家介紹Python利用pandas處理CSV文件的用法示例,文中通過代碼和圖文講解的非常詳細,需要的朋友可以參考下
    2024-07-07
  • Python實現(xiàn)高分辨率圖像導航的代碼

    Python實現(xiàn)高分辨率圖像導航的代碼

    高分辨率圖像導航是一種技術(shù),它允許用戶在大型圖像中進行導航和瀏覽,而無需加載整個圖像到內(nèi)存中,在本文中,我們將使用30行Python代碼實現(xiàn)這一功能,我們將使用Python的圖像處理庫和計算機視覺庫來加載圖像數(shù)據(jù)并生成高分辨率圖像導航
    2024-03-03

最新評論

南京市| 翁源县| 易门县| 清丰县| 荃湾区| 沙湾县| 仁寿县| 舞钢市| 蛟河市| 柳江县| 牡丹江市| 和硕县| 会宁县| 颍上县| 大荔县| 芦溪县| 彭州市| 平乐县| 洛南县| 友谊县| 通渭县| 屯昌县| 长治县| 延安市| 康定县| 蓬安县| 平湖市| 新巴尔虎右旗| 垣曲县| 洛隆县| 临洮县| 缙云县| 禹州市| 长宁区| 湟源县| 镇江市| 左云县| 三台县| 蛟河市| 新昌县| 会东县|