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

Pytorch?PyG實(shí)現(xiàn)EdgePool圖分類(lèi)

 更新時(shí)間:2023年04月21日 09:55:27   作者:實(shí)力  
這篇文章主要為大家介紹了Pytorch?PyG實(shí)現(xiàn)EdgePool圖分類(lèi)示例詳解,有需要的朋友可以借鑒參考下,希望能夠有所幫助,祝大家多多進(jìn)步,早日升職加薪

EdgePool簡(jiǎn)介

EdgePool是一種用于圖分類(lèi)的卷積神經(jīng)網(wǎng)絡(luò)(Convolutional Neural Network,CNN)模型。其主要思想是通過(guò) edge pooling 上下采樣優(yōu)化圖像大小,減少空間復(fù)雜度,提高分類(lèi)性能。

實(shí)現(xiàn)步驟

 數(shù)據(jù)準(zhǔn)備

一般來(lái)講,在構(gòu)建較大規(guī)模數(shù)據(jù)集時(shí),我們都需要對(duì)數(shù)據(jù)進(jìn)行規(guī)范、歸一和清洗處理,以便后續(xù)語(yǔ)義分析或深度學(xué)習(xí)操作。而在圖像數(shù)據(jù)集中,則需使用特定的框架或工具庫(kù)完成。

# 導(dǎo)入MNIST數(shù)據(jù)集
from torch_geometric.datasets import MNISTSuperpixels
# 加載數(shù)據(jù)、劃分訓(xùn)練集和測(cè)試集
dataset = MNISTSuperpixels(root='./mnist', transform=Compose([ToTensor(), NormalizeMeanStd()]))
data = dataset[0]
# 定義超級(jí)參數(shù)
num_features = dataset.num_features
num_classes = dataset.num_classes
# 構(gòu)建訓(xùn)練集和測(cè)試集索引文件
train_mask = torch.zeros(data.num_nodes, dtype=torch.uint8)
train_mask[:60000] = 1
test_mask = torch.zeros(data.num_nodes, dtype=torch.uint8)
test_mask[60000:] = 1
# 創(chuàng)建數(shù)據(jù)加載器
train_loader = DataLoader(data[train_mask], batch_size=32, shuffle=True)
test_loader = DataLoader(data[test_mask], batch_size=32, shuffle=False)

實(shí)現(xiàn)模型

在定義EdgePool模型時(shí),我們需要重新考慮網(wǎng)絡(luò)結(jié)構(gòu)中的上下采樣操作,以便讓整個(gè)網(wǎng)絡(luò)擁有更強(qiáng)大的表達(dá)能力,從而學(xué)習(xí)到更復(fù)雜的關(guān)系。

from torch.nn import Linear
from torch_geometric.nn import EdgePooling
class EdgePool(torch.nn.Module):
    def __init__(self, dataset):
        super(EdgePool, self).__init__()
        # 定義輸入與輸出維度數(shù)
        self.input_dim = dataset.num_features
        self.hidden_dim = 128
        self.output_dim = 10
        # 定義卷積層、歸一化層和pooling層等
        self.conv1 = GCNConv(self.input_dim, self.hidden_dim)
        self.norm1 = BatchNorm1d(self.hidden_dim)
        self.pool1 = EdgePooling(self.hidden_dim)
        self.conv2 = GCNConv(self.hidden_dim, self.hidden_dim)
        self.norm2 = BatchNorm1d(self.hidden_dim)
        self.pool2 = EdgePooling(self.hidden_dim)
        self.conv3 = GCNConv(self.hidden_dim, self.hidden_dim)
        self.norm3 = BatchNorm1d(self.hidden_dim)
        self.pool3 = EdgePooling(self.hidden_dim)
        self.lin = torch.nn.Linear(self.hidden_dim, self.output_dim)
    def forward(self, x, edge_index, batch):
        x = F.relu(self.norm1(self.conv1(x, edge_index)))
        x, edge_index, _, batch, _ = self.pool1(x, edge_index, None, batch)
        x = F.relu(self.norm2(self.conv2(x, edge_index)))
        x, edge_index, _, batch, _ = self.pool2(x, edge_index, None, batch)
        x = F.relu(self.norm3(self.conv3(x, edge_index)))
        x, edge_index, _, batch, _ = self.pool3(x, edge_index, None, batch)
        x = global_mean_pool(x, batch)
        x = self.lin(x)
        return x

在上述代碼中,我們使用了不同的卷積層、池化層和全連接層等神經(jīng)網(wǎng)絡(luò)功能塊來(lái)構(gòu)建EdgePool模型。其中,每個(gè) GCNConv 層被保持為128的隱藏尺寸;BatchNorm1d是一種旨在提高收斂速度并增強(qiáng)網(wǎng)絡(luò)泛化能力的方法;EdgePooling是一種在 GraphConvolution 上附加的特殊類(lèi)別,它將給定圖下采樣至其一半的大小,并返回縮小后的圖與兩個(gè)跟蹤full-graph-to-pool雙向映射(keep and senders)的 edge index(edgendarcs)。 在這種情況下傳遞 None ,表明 batch 未更改。

模型訓(xùn)練

在定義好 EdgePool 網(wǎng)絡(luò)結(jié)構(gòu)之后,需要指定合適的優(yōu)化器、損失函數(shù),并控制訓(xùn)練輪數(shù)、批量大小與學(xué)習(xí)率等超參數(shù)。同時(shí)還要記錄大量日志信息,方便后期跟蹤和駕駛員。

# 定義訓(xùn)練計(jì)劃,包括損失函數(shù)、優(yōu)化器及迭代次數(shù)等
train_epochs = 50
learning_rate = 0.01
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(edge_pool.parameters(), lr=learning_rate)
losses_per_epoch = []
accuracies_per_epoch = []
for epoch in range(train_epochs):
    running_loss = 0.0
    running_corrects = 0.0
    count = 0.0
    for samples in train_loader:
        optimizer.zero_grad()
        x, edge_index, batch = samples.x, samples.edge_index, samples.batch
        out = edge_pool(x, edge_index, batch)
        label = samples.y
        loss = criterion(out, label)
        loss.backward()
        optimizer.step()
        running_loss += loss.item() / len(train_loader.dataset)
        pred = out.argmax(dim=1)
        running_corrects += pred.eq(label).sum().item() / len(train_loader.dataset)
        count += 1
    losses_per_epoch.append(running_loss)
    accuracies_per_epoch.append(running_corrects)
    if (epoch + 1) % 10 == 0:
        print("Train Epoch {}/{} Loss {:.4f} Accuracy {:.4f}".format(
            epoch + 1, train_epochs, running_loss, running_corrects))

在訓(xùn)練過(guò)程中,我們遍歷了每個(gè)批次的數(shù)據(jù),并通過(guò)反向傳播算法進(jìn)行優(yōu)化,并更新了 loss 和 accuracy 輸出值。 同時(shí)方便可視化與記錄,需要將訓(xùn)練過(guò)程中的 loss 和 accuracy 輸出到相應(yīng)的容器中,以便后期進(jìn)行分析和處理。

以上就是Pytorch PyG實(shí)現(xiàn)EdgePool圖分類(lèi)的詳細(xì)內(nèi)容,更多關(guān)于Pytorch PyG EdgePool圖分類(lèi)的資料請(qǐng)關(guān)注腳本之家其它相關(guān)文章!

相關(guān)文章

  • Python filter()及reduce()函數(shù)使用方法解析

    Python filter()及reduce()函數(shù)使用方法解析

    這篇文章主要介紹了Python filter()及reduce()函數(shù)使用方法解析,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下
    2020-09-09
  • PyTorch中的squeeze()和unsqueeze()解析與應(yīng)用案例

    PyTorch中的squeeze()和unsqueeze()解析與應(yīng)用案例

    這篇文章主要介紹了PyTorch中的squeeze()和unsqueeze()解析與應(yīng)用案例,文章內(nèi)容介紹詳細(xì),需要的小伙伴可以參考一下,希望對(duì)你有所幫助
    2022-03-03
  • 用?Python?腳本實(shí)現(xiàn)電腦喚醒后自動(dòng)拍照并截屏發(fā)郵件通知

    用?Python?腳本實(shí)現(xiàn)電腦喚醒后自動(dòng)拍照并截屏發(fā)郵件通知

    這篇文章主要介紹了用?Python?腳本實(shí)現(xiàn)電腦喚醒后自動(dòng)拍照并截屏發(fā)郵件通知,文中詳細(xì)的介紹了代碼示例,具有一定的 參考價(jià)值,感興趣的可以了解一下
    2023-03-03
  • Django CBV類(lèi)的用法詳解

    Django CBV類(lèi)的用法詳解

    這篇文章主要介紹了Django CBV類(lèi)的用法詳解,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友可以參考下
    2019-07-07
  • python生成二維矩陣的兩種方法小結(jié)

    python生成二維矩陣的兩種方法小結(jié)

    本文主要介紹了python生成二維矩陣,包含列表生成m行n列的矩陣和numpy生成想要維度的矩陣的兩種方法,具有一定的參考價(jià)值,感興趣的可以了解一下
    2024-08-08
  • pandas 使用apply同時(shí)處理兩列數(shù)據(jù)的方法

    pandas 使用apply同時(shí)處理兩列數(shù)據(jù)的方法

    下面小編就為大家分享一篇pandas 使用apply同時(shí)處理兩列數(shù)據(jù)的方法,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧
    2018-04-04
  • 淺談tensorflow中張量的提取值和賦值

    淺談tensorflow中張量的提取值和賦值

    今天小編就為大家分享一篇淺談tensorflow中張量的提取值和賦值,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧
    2020-01-01
  • 淺談python中的getattr函數(shù) hasattr函數(shù)

    淺談python中的getattr函數(shù) hasattr函數(shù)

    下面小編就為大家?guī)?lái)一篇淺談python中的getattr函數(shù) hasattr函數(shù)。小編覺(jué)得挺不錯(cuò)的,現(xiàn)在就分享給大家,也給大家做個(gè)參考。一起跟隨小編過(guò)來(lái)看看吧
    2016-06-06
  • 在pycharm中創(chuàng)建django項(xiàng)目的示例代碼

    在pycharm中創(chuàng)建django項(xiàng)目的示例代碼

    這篇文章主要介紹了在pycharm中創(chuàng)建django項(xiàng)目的示例代碼,文中通過(guò)示例代碼介紹的非常詳細(xì),對(duì)大家的學(xué)習(xí)或者工作具有一定的參考學(xué)習(xí)價(jià)值,需要的朋友們下面隨著小編來(lái)一起學(xué)習(xí)學(xué)習(xí)吧
    2020-05-05
  • pytorch 彩色圖像轉(zhuǎn)灰度圖像實(shí)例

    pytorch 彩色圖像轉(zhuǎn)灰度圖像實(shí)例

    今天小編就為大家分享一篇pytorch 彩色圖像轉(zhuǎn)灰度圖像實(shí)例,具有很好的參考價(jià)值,希望對(duì)大家有所幫助。一起跟隨小編過(guò)來(lái)看看吧
    2020-01-01

最新評(píng)論

安阳市| 云和县| 九江县| 佛冈县| 满洲里市| 琼结县| 昌图县| 河池市| 乌鲁木齐市| 繁昌县| 开化县| 仙游县| 夏津县| 金寨县| 定安县| 淮安市| 濮阳市| 林甸县| 齐齐哈尔市| 庆城县| 老河口市| 娱乐| 通化县| 汨罗市| 社会| 那坡县| 永靖县| 界首市| 巴青县| 宝丰县| 新邵县| 桂东县| 雅安市| 陇南市| 平度市| 固原市| 石景山区| 彩票| 乌兰浩特市| 德江县| 抚宁县|