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

pytorch自定義loss損失函數(shù)

 更新時間:2022年02月11日 11:20:32   作者:呆萌的代Ma  
這篇文章主要介紹了pytorch自定義loss損失函數(shù),自定義loss的方法有很多,本文要介紹的是把loss作為一個pytorch的模塊,下面詳細(xì)資料需要的小伙伴可以參考一下

自定義loss的方法有很多,但是在博主查資料的時候發(fā)現(xiàn)有挺多寫法會有問題,靠譜一點的方法是把loss作為一個pytorch的模塊,

比如:

class CustomLoss(nn.Module): # 注意繼承 nn.Module
? ? def __init__(self):
? ? ? ? super(CustomLoss, self).__init__()

? ? def forward(self, x, y):
? ? ? ? # .....這里寫x與y的處理邏輯,即loss的計算方法
? ? ? ? return loss # 注意最后只能返回Tensor值,且?guī)荻?,?loss.requires_grad == True

示例代碼:

以一個pytorch求解線性回歸的代碼為例:

import torch
import torch.nn as nn
import numpy as np
import os

os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"


def get_x_y():
? ? np.random.seed(0)
? ? x = np.random.randint(0, 50, 300)
? ? y_values = 2 * x + 21
? ? x = np.array(x, dtype=np.float32)
? ? y = np.array(y_values, dtype=np.float32)
? ? x = x.reshape(-1, 1)
? ? y = y.reshape(-1, 1)
? ? return x, y


class LinearRegressionModel(nn.Module):
? ? def __init__(self, input_dim, output_dim):
? ? ? ? super(LinearRegressionModel, self).__init__()
? ? ? ? self.linear = nn.Linear(input_dim, output_dim) ?# 輸入的個數(shù),輸出的個數(shù)

? ? def forward(self, x):
? ? ? ? out = self.linear(x)
? ? ? ? return out


if __name__ == '__main__':
? ? input_dim = 1
? ? output_dim = 1
? ? x_train, y_train = get_x_y()

? ? model = LinearRegressionModel(input_dim, output_dim)
? ? epochs = 1000 ?# 迭代次數(shù)
? ? optimizer = torch.optim.SGD(model.parameters(), lr=0.001)
? ? model_loss = nn.MSELoss() # 使用MSE作為loss
? ? # 開始訓(xùn)練模型
? ? for epoch in range(epochs):
? ? ? ? epoch += 1
? ? ? ? # 注意轉(zhuǎn)行成tensor
? ? ? ? inputs = torch.from_numpy(x_train)
? ? ? ? labels = torch.from_numpy(y_train)
? ? ? ? # 梯度要清零每一次迭代
? ? ? ? optimizer.zero_grad()
? ? ? ? # 前向傳播
? ? ? ? outputs: torch.Tensor = model(inputs)
? ? ? ? # 計算損失
? ? ? ? loss = model_loss(outputs, labels)
? ? ? ? # 返向傳播
? ? ? ? loss.backward()
? ? ? ? # 更新權(quán)重參數(shù)
? ? ? ? optimizer.step()
? ? ? ? if epoch % 50 == 0:
? ? ? ? ? ? print('epoch {}, loss {}'.format(epoch, loss.item()))

步驟1:添加自定義的類

我們就用自定義的寫法來寫與MSE相同的效果,MSE計算公式如下:

添加一個類:

class CustomLoss(nn.Module):
? ? def __init__(self):
? ? ? ? super(CustomLoss, self).__init__()
? ? ? ? self.mse_loss = nn.MSELoss()

? ? def forward(self, x, y):
? ? ? ? mse_loss = torch.mean(torch.pow((x - y), 2)) # x與y相減后平方,求均值即為MSE
? ? ? ? return mse_loss

步驟2:修改使用的loss函數(shù)

只需要把原始代碼中的:

model_loss = nn.MSELoss() # 使用MSE作為loss

改為:

model_loss = CustomLoss() ?# 自定義loss

即可

完整代碼:

import torch
import torch.nn as nn
import numpy as np
import os

os.environ["KMP_DUPLICATE_LIB_OK"] = "TRUE"


def get_x_y():
? ? np.random.seed(0)
? ? x = np.random.randint(0, 50, 300)
? ? y_values = 2 * x + 21
? ? x = np.array(x, dtype=np.float32)
? ? y = np.array(y_values, dtype=np.float32)
? ? x = x.reshape(-1, 1)
? ? y = y.reshape(-1, 1)
? ? return x, y


class LinearRegressionModel(nn.Module):
? ? def __init__(self, input_dim, output_dim):
? ? ? ? super(LinearRegressionModel, self).__init__()
? ? ? ? self.linear = nn.Linear(input_dim, output_dim) ?# 輸入的個數(shù),輸出的個數(shù)

? ? def forward(self, x):
? ? ? ? out = self.linear(x)
? ? ? ? return out


class CustomLoss(nn.Module):
? ? def __init__(self):
? ? ? ? super(CustomLoss, self).__init__()
? ? ? ? self.mse_loss = nn.MSELoss()

? ? def forward(self, x, y):
? ? ? ? mse_loss = torch.mean(torch.pow((x - y), 2))
? ? ? ? return mse_loss


if __name__ == '__main__':
? ? input_dim = 1
? ? output_dim = 1
? ? x_train, y_train = get_x_y()

? ? model = LinearRegressionModel(input_dim, output_dim)
? ? epochs = 1000 ?# 迭代次數(shù)
? ? optimizer = torch.optim.SGD(model.parameters(), lr=0.001)
? ? # model_loss = nn.MSELoss() # 使用MSE作為loss
? ? model_loss = CustomLoss() ?# 自定義loss
? ? # 開始訓(xùn)練模型
? ? for epoch in range(epochs):
? ? ? ? epoch += 1
? ? ? ? # 注意轉(zhuǎn)行成tensor
? ? ? ? inputs = torch.from_numpy(x_train)
? ? ? ? labels = torch.from_numpy(y_train)
? ? ? ? # 梯度要清零每一次迭代
? ? ? ? optimizer.zero_grad()
? ? ? ? # 前向傳播
? ? ? ? outputs: torch.Tensor = model(inputs)
? ? ? ? # 計算損失
? ? ? ? loss = model_loss(outputs, labels)
? ? ? ? # 返向傳播
? ? ? ? loss.backward()
? ? ? ? # 更新權(quán)重參數(shù)
? ? ? ? optimizer.step()
? ? ? ? if epoch % 50 == 0:
? ? ? ? ? ? print('epoch {}, loss {}'.format(epoch, loss.item()))

到此這篇關(guān)于pytorch自定義loss損失函數(shù)的文章就介紹到這了,更多相關(guān)pytorch loss損失函數(shù)內(nèi)容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關(guān)文章希望大家以后多多支持腳本之家!

相關(guān)文章

  • pytest用例執(zhí)行順序和跳過執(zhí)行詳解

    pytest用例執(zhí)行順序和跳過執(zhí)行詳解

    本文主要介紹了pytest用例執(zhí)行順序和跳過執(zhí)行詳解,文中通過示例代碼介紹的非常詳細(xì),對大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2023-02-02
  • pyecharts結(jié)合flask框架的使用

    pyecharts結(jié)合flask框架的使用

    這篇文章主要介紹了pyecharts結(jié)合flask框架,主要是介紹如何在Flask框架中使用pyecharts,本文通過示例代碼給大家介紹的非常詳細(xì),需要的朋友可以參考下
    2022-06-06
  • python析構(gòu)函數(shù)用法及注意事項

    python析構(gòu)函數(shù)用法及注意事項

    在本篇文章里小編給大家整理的是一篇關(guān)于python析構(gòu)函數(shù)用法及注意事項,有需要的朋友們可以學(xué)習(xí)參考下。
    2021-06-06
  • Python?gRPC流式通信協(xié)議詳細(xì)講解

    Python?gRPC流式通信協(xié)議詳細(xì)講解

    這篇文章主要介紹了Python?gRPC流式通信協(xié)議,最近幾天在搞golang的grpc,跑通之后想用php作為客戶端調(diào)用一下grpc服務(wù),結(jié)果拉了,一個php的grpc服務(wù)安裝,搞了好幾天,總算搞定了
    2022-11-11
  • Python3.6筆記之將程序運行結(jié)果輸出到文件的方法

    Python3.6筆記之將程序運行結(jié)果輸出到文件的方法

    下面小編就為大家分享一篇Python3.6筆記之將程序運行結(jié)果輸出到文件的方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2018-04-04
  • PyTorch加載模型model.load_state_dict()問題及解決

    PyTorch加載模型model.load_state_dict()問題及解決

    這篇文章主要介紹了PyTorch加載模型model.load_state_dict()問題及解決,具有很好的參考價值,希望對大家有所幫助。
    2023-02-02
  • 利用python將圖片轉(zhuǎn)換成excel文檔格式

    利用python將圖片轉(zhuǎn)換成excel文檔格式

    編寫了一小段Python代碼,將圖片轉(zhuǎn)為了Excel,純屬娛樂,下面這篇文章主要給大家介紹了關(guān)于利用python將圖片轉(zhuǎn)換成excel文檔格式的相關(guān)資料,需要的朋友可以參考借鑒,下面來一起看看吧。
    2017-12-12
  • Python裝飾器使用你可能不知道的幾種姿勢

    Python裝飾器使用你可能不知道的幾種姿勢

    這篇文章主要給大家介紹了關(guān)于Python裝飾器使用你可能不知道的幾種姿勢,文中通過示例代碼介紹的非常詳細(xì),對大家的學(xué)習(xí)或者使用Python具有一定的參考學(xué)習(xí)價值,需要的朋友們下面來一起學(xué)習(xí)學(xué)習(xí)吧
    2019-10-10
  • python生成單位陣或?qū)顷嚨娜N方式小結(jié)

    python生成單位陣或?qū)顷嚨娜N方式小結(jié)

    這篇文章主要介紹了python生成單位陣或?qū)顷嚨娜N方式小結(jié),具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教
    2022-05-05
  • python中hasattr()、getattr()、setattr()函數(shù)的使用

    python中hasattr()、getattr()、setattr()函數(shù)的使用

    這篇文章主要介紹了python中hasattr()、getattr()、setattr()函數(shù)的使用方法,本文給大家介紹的非常詳細(xì),具有一定的參考借鑒價值,需要的朋友可以參考下
    2019-08-08

最新評論

丹江口市| 保山市| 连江县| 大荔县| 南澳县| 太谷县| 厦门市| 喀喇沁旗| 镇原县| 龙岩市| 怀化市| 白沙| 衢州市| 泰安市| 客服| 芦溪县| 翁源县| 南靖县| 巢湖市| 泾源县| 长汀县| 琼中| 隆子县| 镇原县| 潜江市| 沁阳市| 利川市| 北辰区| 叶城县| 崇州市| 新密市| 桦甸市| 博乐市| 凤阳县| 大足县| 淮滨县| 香格里拉县| 郸城县| 武威市| 含山县| 上林县|