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

詳解Pytorch+PyG實現(xiàn)GAT過程示例

 更新時間:2023年04月21日 10:01:31   作者:實力  
這篇文章主要為大家介紹了Pytorch+PyG實現(xiàn)GAT過程示例詳解,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪

導(dǎo)入庫和數(shù)據(jù)

GAT(圖注意力網(wǎng)絡(luò))是常見的圖神經(jīng)網(wǎng)絡(luò)結(jié)構(gòu)之一,它使用注意力機制來對節(jié)點進(jìn)行特征加權(quán),并考慮其鄰居節(jié)點的交互。

首先,我們需要導(dǎo)入PyTorch和PyG庫,然后準(zhǔn)備好我們的數(shù)據(jù)。例如,我們可以使用以下方式生成一個簡單的隨機數(shù)據(jù)集:

from torch_geometric.datasets import Planetoid
dataset = Planetoid(root='/tmp/Cora', name='Cora')
train_loader = DataLoader(dataset[0], batch_size=128, shuffle=True)
test_loader = DataLoader(dataset[0], batch_size=128, shuffle=False)

其中, Planetoid 是PyG提供的圖形數(shù)據(jù)集之一。這里我們選擇了 Cora 數(shù)據(jù)集并存儲到 /tmp/Cora 文件夾中。然后我們將該數(shù)據(jù)集分成訓(xùn)練集和測試集,設(shè)置相應(yīng)的加載器。

定義模型結(jié)構(gòu)

接下來,我們需要定義GAT模型的結(jié)構(gòu)。通過PyTorch和PyG,我們可以自己定義完整的GAT模型或者利用現(xiàn)有的庫函數(shù)快速構(gòu)建模型。在這里,我們將使用 torch_geometric.nn.GATConv 函數(shù)逐層堆疊多個圖注意力層來實現(xiàn)GAT模型。以下是GAT模型定義的示例代碼:

import torch.nn.functional as F
from torch_geometric.nn import GATConv
class Net(torch.nn.Module):
    def __init__(self, in_channels, out_channels):
        super(Net, self).__init__()
        self.num_layers = 2
        self.conv1 = GATConv(in_channels=in_channels, out_channels=16, heads=8, dropout=0.6)
        self.conv2 = GATConv(in_channels=16*8, out_channels=out_channels, heads=1, concat=False, dropout=0.6)
    def forward(self, data):
        x, edge_index = data.x, data.edge_index
        x = F.dropout(x, p=0.6, training=self.training)
        x = F.elu(self.conv1(x, edge_index))
        x = F.dropout(x, p=0.6, training=self.training)
        x = self.conv2(x, edge_index)
        return F.log_softmax(x, dim=1)

上述代碼中,我們定義了一個 Net 類用于構(gòu)建GAT網(wǎng)絡(luò),接收輸入通道數(shù)和輸出通道數(shù)作為參數(shù)。例如,我們可以按照以下方式創(chuàng)建一個將 CORD 參量作為輸入特征向量大小、64 個隱藏節(jié)點(每個注意力頭)。并將數(shù)字類別作為輸出大小的GAT模型:

model = Net(in_channels=dataset.num_features, out_channels=dataset.num_classes)

其中 num_featuresnum_classes 是PyG數(shù)據(jù)集中包含的屬性。

定義訓(xùn)練函數(shù)

然后,我們需要定義訓(xùn)練函數(shù)來訓(xùn)練我們的GAT神經(jīng)網(wǎng)絡(luò)。在這里,我們將使用交叉熵?fù)p失和Adam優(yōu)化器進(jìn)行訓(xùn)練,并在每一個epoch結(jié)束時計算準(zhǔn)確率并打印出來。以下是訓(xùn)練函數(shù)的示例代碼:

import torch.optim as optim
from tqdm import tqdm
def train(model, loader, optimizer, loss_fn):
    model.train()
    correct = 0
    total_loss = 0
    for data in tqdm(loader, desc='Training'):
        optimizer.zero_grad()
        out = model(data)
        pred = out.argmax(dim=1)
        loss = loss_fn(out[data.train_mask], data.y[data.train_mask])
        loss.backward()
        optimizer.step()
        total_loss += loss.item() * data.num_graphs
        correct += pred[data.train_mask].eq(data.y[data.train_mask]).sum().item()
    return total_loss / len(loader.dataset), correct / len(data.train_mask)

在上述代碼中,我們遍歷加載器中的每個數(shù)據(jù)批次,并對模型進(jìn)行培訓(xùn)。對于每個圖數(shù)據(jù)批次,我們計算網(wǎng)絡(luò)輸出、預(yù)測和損失,然后通過反向傳播來更新權(quán)重。最后,我們將總損失和正確率記錄下來并返回。

定義測試函數(shù)

接下來,我們還需要定義測試函數(shù)來測試我們的GAT神經(jīng)網(wǎng)絡(luò)性能表現(xiàn)。我們將利用與訓(xùn)練函數(shù)相同的輸出參數(shù)進(jìn)行測試,并打印出最終的測試準(zhǔn)確率。以下是測試函數(shù)的示例代碼:

def test(model, loader, loss_fn):
    model.eval()
    correct = 0
    total_loss = 0
    with torch.no_grad():
        for data in tqdm(loader, desc='Testing'):
            out = model(data)
            pred = out.argmax(dim=1)
            loss = loss_fn(out[data.test_mask], data.y[data.test_mask])
            total_loss += loss.item() * data.num_graphs
            correct += pred[data.test_mask].eq(data.y[data.test_mask]).sum().item()
    return total_loss / len(loader.dataset), correct / len(data.test_mask)

在上述代碼中,我們對測試數(shù)據(jù)集中的所有數(shù)據(jù)進(jìn)行了循環(huán),并計算網(wǎng)絡(luò)的輸出和預(yù)測。我們記錄下總損失和正確分類的數(shù)據(jù)量,并返回?fù)p失和準(zhǔn)確率之間的比率。

訓(xùn)練模型并評估訓(xùn)練結(jié)果

最后,我們可以使用前面定義過的函數(shù)來定義主函數(shù),從而完成GAT神經(jīng)網(wǎng)絡(luò)的訓(xùn)練和測試。以下是主函數(shù)的示例代碼:

if __name__ == '__main__':
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    model = Net(in_channels=dataset.num_features, out_channels=dataset.num_classes).to(device)
    train_loader = DataLoader(dataset[0], batch_size=128, shuffle=True)
    test_loader = DataLoader(dataset[0], batch_size=128, shuffle=False)
    optimizer = optim.Adam(model.parameters(), lr=0.01)
    loss_fn = nn.CrossEntropyLoss()
    for epoch in range(1, 201):
        train_loss, train_acc = train(model, train_loader, optimizer, loss_fn)
        test_loss, test_acc = test(model, test_loader, loss_fn)
        print(f'Epoch {epoch:03d}, Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, '
              f'Test Loss: {test_loss:.4f}, Test Acc: {test_acc:.4f}')

通過上述代碼,我們就可以完成GAT神經(jīng)網(wǎng)絡(luò)的訓(xùn)練和測試。我們使用 DataLoader 函數(shù)進(jìn)行數(shù)據(jù)加載,設(shè)置學(xué)習(xí)率、損失函數(shù)、訓(xùn)練輪數(shù)等超參數(shù)。最后,我們可以在屏幕上看到每個時代的準(zhǔn)確率和損失值,并通過它們評估模型的訓(xùn)練表現(xiàn)。

以上就是詳解Pytorch+PyG實現(xiàn)GAT過程示例的詳細(xì)內(nèi)容,更多關(guān)于Pytorch PyG實現(xiàn)GAT的資料請關(guān)注腳本之家其它相關(guān)文章!

相關(guān)文章

  • python實現(xiàn)根據(jù)主機名字獲得所有ip地址的方法

    python實現(xiàn)根據(jù)主機名字獲得所有ip地址的方法

    這篇文章主要介紹了python實現(xiàn)根據(jù)主機名字獲得所有ip地址的方法,涉及Python解析IP地址的相關(guān)技巧,需要的朋友可以參考下
    2015-06-06
  • python 實現(xiàn)將文件或文件夾用相對路徑打包為 tar.gz 文件的方法

    python 實現(xiàn)將文件或文件夾用相對路徑打包為 tar.gz 文件的方法

    今天小編就為大家分享一篇python 實現(xiàn)將文件或文件夾用相對路徑打包為 tar.gz 文件的方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2019-06-06
  • python實現(xiàn)防截圖的6種方法詳解

    python實現(xiàn)防截圖的6種方法詳解

    防截圖是指一組技術(shù)或方法,用于防止他人在未經(jīng)允許的情況下在屏幕上截取或記錄圖像,這是一個重要的安全措施,它可以防止竊取敏感信息或監(jiān)視個人信息,本文為大家整理了6種python可以防截圖的方法,需要的可以參考下
    2023-10-10
  • python跨文件夾調(diào)用別的文件夾下py文件或參數(shù)方式詳解

    python跨文件夾調(diào)用別的文件夾下py文件或參數(shù)方式詳解

    這篇文章主要給大家介紹了關(guān)于python跨文件夾調(diào)用別的文件夾下py文件或參數(shù)方式的相關(guān)資料,在python中有時候我們需要調(diào)用另一.py文件中的方法或者類,需要的朋友可以參考下
    2023-08-08
  • python 異常捕獲詳解流程

    python 異常捕獲詳解流程

    異常即非正常狀態(tài),在Python中使用異常對象來表示異常。若程序在編譯或運行過程中發(fā)生錯誤,程序的執(zhí)行過程就會發(fā)生改變,拋出異常對象,程序流進(jìn)入異常處理。如果異常對象沒有被處理或捕捉,程序就會執(zhí)行回溯(Traceback)來終止程序
    2022-03-03
  • 詳解duck typing鴨子類型程序設(shè)計與Python的實現(xiàn)示例

    詳解duck typing鴨子類型程序設(shè)計與Python的實現(xiàn)示例

    這篇文章主要介紹了詳解duck typing鴨子類型程序設(shè)計與Python的實現(xiàn)示例,鴨子類型特指解釋型語言中的一種編程風(fēng)格,需要的朋友可以參考下
    2016-06-06
  • python正則表達(dá)式對字符串的查找匹配

    python正則表達(dá)式對字符串的查找匹配

    正則表達(dá)式是一種文本模式,包括普通字符(例如,a 到 z 之間的字母)和特殊字符(稱為“元字符”),下面這篇文章主要給大家介紹了關(guān)于python正則表達(dá)式對字符串的查找匹配的相關(guān)資料,需要的朋友可以參考下
    2022-09-09
  • python中循環(huán)語句while用法實例

    python中循環(huán)語句while用法實例

    這篇文章主要介紹了python中循環(huán)語句while用法,實例分析了while語句的使用方法,需要的朋友可以參考下
    2015-05-05
  • Python實現(xiàn)的最近最少使用算法

    Python實現(xiàn)的最近最少使用算法

    這篇文章主要介紹了Python實現(xiàn)的最近最少使用算法,涉及節(jié)點、時間、流程控制等相關(guān)技巧,需要的朋友可以參考下
    2015-07-07
  • Python實現(xiàn)爬取逐浪小說的方法

    Python實現(xiàn)爬取逐浪小說的方法

    這篇文章主要介紹了Python實現(xiàn)爬取逐浪小說的方法,基于Python的正則匹配功能實現(xiàn)爬取小說頁面標(biāo)題、鏈接及正文等功能,需要的朋友可以參考下
    2015-07-07

最新評論

泗洪县| 浠水县| 乌鲁木齐县| 司法| 万源市| 林甸县| 且末县| 泰兴市| 游戏| 敦化市| 汉源县| 余干县| 天等县| 余干县| 德昌县| 贡山| 县级市| 陆丰市| 克拉玛依市| 张家界市| 邻水| 广西| 杨浦区| 东台市| 新民市| 宁远县| 宣汉县| 荥经县| 吉林市| 辽源市| 连州市| 慈利县| 隆德县| 贵阳市| 昆明市| 顺昌县| 榕江县| 中超| 康保县| 鄂尔多斯市| 德清县|