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

PyG搭建GCN模型實現(xiàn)節(jié)點分類GCNConv參數(shù)詳解

 更新時間:2022年05月10日 15:23:46   作者:Cyril_KI  
這篇文章主要為大家介紹了PyG搭建GCN模型實現(xiàn)節(jié)點分類GCNConv參數(shù)詳解,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪

前言

在上一篇文章PyG搭建GCN前的準(zhǔn)備:了解PyG中的數(shù)據(jù)格式中,大致了解了PyG中的數(shù)據(jù)格式,這篇文章主要是簡單搭建GCN來實現(xiàn)節(jié)點分類,主要目的是了解PyG中GCN的參數(shù)情況。

模型搭建

首先導(dǎo)入包:

from torch_geometric.nn import GCNConv

模型參數(shù):

in_channels:輸入通道,比如節(jié)點分類中表示每個節(jié)點的特征數(shù)。

out_channels:輸出通道,最后一層GCNConv的輸出通道為節(jié)點類別數(shù)(節(jié)點分類)。

improved:如果為True表示自環(huán)增加,也就是原始鄰接矩陣加上2I而不是I,默認(rèn)為False。

cached:如果為True,GCNConv在第一次對鄰接矩陣進(jìn)行歸一化時會進(jìn)行緩存,以后將不再重復(fù)計算。

add_self_loops:如果為False不再強(qiáng)制添加自環(huán),默認(rèn)為True。

normalize:默認(rèn)為True,表示對鄰接矩陣進(jìn)行歸一化。

bias:默認(rèn)添加偏置。

于是模型搭建如下:

class GCN(torch.nn.Module):
    def __init__(self, num_node_features, num_classes):
        super(GCN, self).__init__()
        self.conv1 = GCNConv(num_node_features, 16)
        self.conv2 = GCNConv(16, num_classes)
    def forward(self, data):
        x, edge_index = data.x, data.edge_index
        x = self.conv1(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, training=self.training)
        x = self.conv2(x, edge_index)
        x = F.relu(x)
        x = F.dropout(x, training=self.training)
        x = F.softmax(x, dim=1)
        return x

輸出一下模型:

data = Planetoid(root='/data/CiteSeer', name='CiteSeer')model = GCN(data.num_node_features, data.num_classes).to(device)print(model)GCN(
  (conv1): GCNConv(3703, 16)
  (conv2): GCNConv(16, 6)
)

輸出為:

GCN( (conv1): GCNConv(3703, 16) (conv2): GCNConv(16, 6))GCN(
  (conv1): GCNConv(3703, 16)
  (conv2): GCNConv(16, 6)
)

1. 前向傳播

查看官方文檔中GCNConv的輸入輸出要求:

可以發(fā)現(xiàn),GCNConv中需要輸入的是節(jié)點特征矩陣x和鄰接關(guān)系edge_index,還有一個可選項edge_weight。因此我們首先:

x, edge_index = data.x, data.edge_index
x = self.conv1(x, edge_index)
x = F.relu(x)
x = F.dropout(x, training=self.training)

此時我們不妨輸出一下x及其size:

tensor([[0.0000, 0.1630, 0.0000,  ..., 0.0000, 0.0488, 0.0000],
        [0.0000, 0.2451, 0.1614,  ..., 0.0000, 0.0125, 0.0000],
        [0.1175, 0.0262, 0.2141,  ..., 0.2592, 0.0000, 0.0000],
        ...,
        [0.0000, 0.0000, 0.0000,  ..., 0.0000, 0.1825, 0.0000],
        [0.0000, 0.1024, 0.0000,  ..., 0.0498, 0.0000, 0.0000],
        [0.0000, 0.3263, 0.0000,  ..., 0.0000, 0.0000, 0.0000]],
       device='cuda:0', grad_fn=<FusedDropoutBackward0>)
torch.Size([3327, 16])

此時的x一共3327行,每一行表示一個節(jié)點經(jīng)過第一層卷積更新后的狀態(tài)向量。

那么同理,由于:

self.conv2 = GCNConv(16, num_classes)

所以經(jīng)過第二層卷積后:

x = self.conv2(x, edge_index)x = F.relu(x)x = F.dropout(x, training=self.training)x = self.conv2(x, edge_index)
x = F.relu(x)
x = F.dropout(x, training=self.training)

此時得到的x的size應(yīng)該為:

torch.Size([3327, 6])

即每個節(jié)點的維度為6的狀態(tài)向量。

由于我們需要進(jìn)行6分類,所以最后需要加上一個softmax:

x = F.softmax(x, dim=1)

dim=1表示對每一行進(jìn)行運算,最終每一行之和加起來為1,也就表示了該節(jié)點為每一類的概率。輸出此時的x:

tensor([[0.1607, 0.1727, 0.1607, 0.1607, 0.1607, 0.1846], [0.1654, 0.1654, 0.1654, 0.1654, 0.1654, 0.1731], [0.1778, 0.1622, 0.1733, 0.1622, 0.1622, 0.1622], ..., [0.1659, 0.1659, 0.1659, 0.1704, 0.1659, 0.1659], [0.1667, 0.1667, 0.1667, 0.1667, 0.1667, 0.1667], [0.1641, 0.1641, 0.1658, 0.1766, 0.1653, 0.1641]], device='cuda:0', grad_fn=<SoftmaxBackward0>)tensor([[0.1607, 0.1727, 0.1607, 0.1607, 0.1607, 0.1846],
        [0.1654, 0.1654, 0.1654, 0.1654, 0.1654, 0.1731],
        [0.1778, 0.1622, 0.1733, 0.1622, 0.1622, 0.1622],
        ...,
        [0.1659, 0.1659, 0.1659, 0.1704, 0.1659, 0.1659],
        [0.1667, 0.1667, 0.1667, 0.1667, 0.1667, 0.1667],
        [0.1641, 0.1641, 0.1658, 0.1766, 0.1653, 0.1641]], device='cuda:0',
       grad_fn=<SoftmaxBackward0>)

2. 反向傳播

在訓(xùn)練時,我們首先利用前向傳播計算出輸出:

out = model(data)

out即為最終得到的每個節(jié)點的6個概率值,但在實際訓(xùn)練中,我們只需要計算出訓(xùn)練集的損失,所以損失函數(shù)這樣寫:

loss = loss_function(out[data.train_mask], data.y[data.train_mask])

然后計算梯度,反向更新!

3. 訓(xùn)練

訓(xùn)練的完整代碼:

def train(): optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4) loss_function = torch.nn.CrossEntropyLoss().to(device) model.train() for epoch in range(500): out = model(data) optimizer.zero_grad() loss = loss_function(out[data.train_mask], data.y[data.train_mask]) loss.backward() optimizer.step() print('Epoch {:03d} loss {:.4f}'.format(epoch, loss.item()))def train():
    optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
    loss_function = torch.nn.CrossEntropyLoss().to(device)
    model.train()
    for epoch in range(500):
        out = model(data)
        optimizer.zero_grad()
        loss = loss_function(out[data.train_mask], data.y[data.train_mask])
        loss.backward()
        optimizer.step()
        print('Epoch {:03d} loss {:.4f}'.format(epoch, loss.item()))

4. 測試

我們首先需要算出模型對所有節(jié)點的預(yù)測值:

model(data)

此時得到的是每個節(jié)點的6個概率值,我們需要在每一行上取其最大值:

model(data).max(dim=1)

輸出一下:

torch.return_types.max(
values=tensor([0.9100, 0.9071, 0.9786,  ..., 0.4321, 0.4009, 0.8779], device='cuda:0',
       grad_fn=<MaxBackward0>),
indices=tensor([3, 1, 5,  ..., 3, 1, 5], device='cuda:0'))

返回的第一項是每一行的最大值,第二項為最大值在這一行中的索引,我們只需要取第二項,那么最終的預(yù)測值應(yīng)該寫為:

_, pred = model(data).max(dim=1)

然后計算預(yù)測精度:

correct = int(pred[data.test_mask].eq(data.y[data.test_mask]).sum().item())
acc = correct / int(data.test_mask.sum())
print('GCN Accuracy: {:.4f}'.format(acc))

完整代碼

完整代碼中實現(xiàn)了論文中提到的四種數(shù)據(jù)集,代碼地址:PyG-GCN。

以上就是PyG搭建GCN模型實現(xiàn)節(jié)點分類GCNConv參數(shù)詳解的詳細(xì)內(nèi)容,更多關(guān)于PyG搭建GCNConv節(jié)點分類的資料請關(guān)注腳本之家其它相關(guān)文章!

相關(guān)文章

  • python3應(yīng)用windows api對后臺程序窗口及桌面截圖并保存的方法

    python3應(yīng)用windows api對后臺程序窗口及桌面截圖并保存的方法

    今天小編就為大家分享一篇python3應(yīng)用windows api對后臺程序窗口及桌面截圖并保存的方法,具有很好的參考價值,希望對大家有所幫助。一起跟隨小編過來看看吧
    2019-08-08
  • Python并發(fā)編程之未來模塊Futures

    Python并發(fā)編程之未來模塊Futures

    這篇文章主要為大家介紹了Python的未來,python并發(fā)編程之未來模塊Futures的詳解,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪
    2022-05-05
  • PyQt5類型判定+對象刪除操作

    PyQt5類型判定+對象刪除操作

    這篇文章主要介紹了PyQt5類型判定+對象刪除操作,本文通過實例代碼給大家介紹的非常詳細(xì),感興趣的朋友跟隨小編一起看看吧
    2024-06-06
  • python驗證碼圖片處理(二值化)

    python驗證碼圖片處理(二值化)

    這篇文章主要介紹了python驗證碼圖片處理(二值化),文中通過示例代碼介紹的非常詳細(xì),對大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2019-11-11
  • Django 通過JS實現(xiàn)ajax過程詳解

    Django 通過JS實現(xiàn)ajax過程詳解

    這篇文章主要介紹了Django 通過JS實現(xiàn)ajax過程詳解,文中通過示例代碼介紹的非常詳細(xì),對大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價值,需要的朋友可以參考下
    2019-07-07
  • 檢查Python中的變量是否是字符串(兩種不同方法)

    檢查Python中的變量是否是字符串(兩種不同方法)

    數(shù)據(jù)類型是編程語言最重要的特征,它區(qū)分了我們可以存儲的不同類型的數(shù)據(jù),如字符串、int和float,這篇文章主要介紹了兩種不同的方法來檢查Python中的變量是否是字符串,需要的朋友可以參考下
    2023-08-08
  • 深入解析Python設(shè)計模式編程中建造者模式的使用

    深入解析Python設(shè)計模式編程中建造者模式的使用

    這篇文章主要介紹了深入解析Python設(shè)計模式編程中建造者模式的使用,建造者模式的程序通常將所有細(xì)節(jié)都交由子類實現(xiàn),需要的朋友可以參考下
    2016-03-03
  • python密碼學(xué)RSA密碼加密教程

    python密碼學(xué)RSA密碼加密教程

    這篇文章主要為大家介紹了python密碼學(xué)RSA密碼加密教程,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪
    2022-05-05
  • 淺談python常用程序算法

    淺談python常用程序算法

    這篇文章主要介紹了python常用程序算法,文中通過示例代碼介紹的非常詳細(xì),對大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價值,需要的朋友們下面隨著小編來一起學(xué)習(xí)學(xué)習(xí)吧
    2019-03-03
  • opencv用VS2013調(diào)試時用Image Watch插件查看圖片

    opencv用VS2013調(diào)試時用Image Watch插件查看圖片

    本文主要介紹了opencv用VS2013調(diào)試時用Image Watch插件查看圖片,直接以圖片形式可視化了opencv中的Mat變量。感興趣的可以了解下
    2021-07-07

最新評論

宣恩县| 红原县| 介休市| 肃南| 崇信县| 嘉善县| 龙南县| 衡山县| 桐庐县| 桐庐县| 襄樊市| 沭阳县| 白山市| 楚雄市| 博爱县| 河西区| 屏南县| 黑水县| 皋兰县| 长葛市| 宁津县| 镇远县| 湖北省| 新野县| 冕宁县| 景东| 南溪县| 宝兴县| 潮州市| 婺源县| 麟游县| 商都县| 沾化县| 武夷山市| 会东县| 资阳市| 南陵县| 尼玛县| 临桂县| 肃北| 诸城市|