GAT学习:PyG实现GAT(自定义GAT层)网络(四)

2024-02-01 08:18

本文主要是介绍GAT学习:PyG实现GAT(自定义GAT层)网络(四),希望对大家解决编程问题提供一定的参考价值,需要的开发者们随着小编来一起学习吧!

PyG实现自定义GAT层

  • 完整代码

本系列中的第三篇介绍了如何调用pyg封装好的GAT函数,当然同样的,我们需要学会如何自定义网络层以满足研究需求。

完整代码

import torch
import math
from torch_geometric.nn import MessagePassing
from torch_geometric.utils import add_self_loops,remove_self_loops,softmax
from torch_geometric.datasets import Planetoid
import ssl
import torch.nn.functional as Fclass GATConv(MessagePassing):def __init__(self, in_channels,out_channels, heads: int = 1, concat: bool = True,negative_slope: float = 0.2, dropout: float = 0.,add_self_loops: bool = True, bias: bool = True, **kwargs):kwargs.setdefault('aggr', 'add')super(GATConv, self).__init__(node_dim=0, **kwargs)#in_channel&out channel就是我们的输入输出数self.in_channels = in_channelsself.out_channels = out_channels#head即设置几个attention头self.heads = heads#concat用于设置是否拼接attention的输出self.concat = concat#negative_slope设置leaklyRelu的参数self.negative_slope = negative_slopeself.dropout = dropout#add_self_loops设置是否添加自环self.add_self_loops = add_self_loops#这里将特征映射到每个attention头所需要的特征数,从而满足每个attention头的输入self.lin = Linear(in_channels, heads * out_channels, bias=False)self.att = Parameter(torch.Tensor(1, heads, out_channels))if bias and concat:self.bias = torch.nn.Parameter(torch.Tensor(heads * out_channels))elif bias and not concat:self.bias = torch.nn.Parameter(torch.Tensor(out_channels))else:self.register_parameter('bias', None)self._alpha = None#用于重置参数self.reset_parameters()def reset_parameters(self):glorot(self.lin.weight)glorot(self.att)zeros(self.bias)def forward(self, x, edge_index, return_attention_weights=None):H, C = self.heads, self.out_channelsx = self.lin(x).view(-1, H, C)#这里alpha的规模为[node_num,heads]alpha = (x * self.att).sum(dim=-1)if self.add_self_loops:num_nodes = x.size(0)num_nodes = x.size(0) if x is not None else num_nodesedge_index, _ = remove_self_loops(edge_index)edge_index, _ = add_self_loops(edge_index, num_nodes=num_nodes)# propagate_type: (x: OptPairTensor, alpha: OptPairTensor)out = self.propagate(edge_index, x=x,alpha=alpha)alpha = self._alphaself._alpha = Noneif self.concat:out = out.view(-1, self.heads * self.out_channels)else:out = out.mean(dim=1)if self.bias is not None:out += self.biasif isinstance(return_attention_weights, bool):return out, (edge_index, alpha)else:return outdef message(self, x_j, alpha_j, index):alpha = alpha_j#alpha_j[edge_num,heads]alpha = F.leaky_relu(alpha, self.negative_slope)alpha = softmax(alpha, index)self._alpha = alphaalpha = F.dropout(alpha, p=self.dropout, training=self.training)return x_j * alpha.unsqueeze(-1)class Net(torch.nn.Module):def __init__(self):super(Net,self).__init__()self.gat1=GATConv(dataset.num_node_features,8,8,dropout=0.6)self.gat2=GATConv(64,7,1,dropout=0.6)def forward(self,data):x,edge_index=data.x, data.edge_indexx=self.gat1(x,edge_index)x=self.gat2(x,edge_index)return F.log_softmax(x,dim=1)dataset = Planetoid(root='Cora', name='Cora')
x=dataset[0].x
edge_index=dataset[0].edge_indexdevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = Net().to(device)
data = dataset[0].to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)model.train()
for epoch in range(100):optimizer.zero_grad()out = model(data)loss = F.nll_loss(out[data.train_mask], data.y[data.train_mask])loss.backward()optimizer.step()model.eval()
_, pred = model(data).max(dim=1)
correct = int(pred[data.test_mask].eq(data.y[data.test_mask]).sum().item())
acc = correct/int(data.test_mask.sum())
print('Accuracy:{:.4f}'.format(acc))
>>>Accuracy:0.7930

这篇关于GAT学习:PyG实现GAT(自定义GAT层)网络(四)的文章就介绍到这儿,希望我们推荐的文章对编程师们有所帮助!



http://www.chinasem.cn/article/666654

相关文章

Java实现优雅日期处理的方案详解

《Java实现优雅日期处理的方案详解》在我们的日常工作中,需要经常处理各种格式,各种类似的的日期或者时间,下面我们就来看看如何使用java处理这样的日期问题吧,感兴趣的小伙伴可以跟随小编一起学习一下... 目录前言一、日期的坑1.1 日期格式化陷阱1.2 时区转换二、优雅方案的进阶之路2.1 线程安全重构2

Android实现两台手机屏幕共享和远程控制功能

《Android实现两台手机屏幕共享和远程控制功能》在远程协助、在线教学、技术支持等多种场景下,实时获得另一部移动设备的屏幕画面,并对其进行操作,具有极高的应用价值,本项目旨在实现两台Android手... 目录一、项目概述二、相关知识2.1 MediaProjection API2.2 Socket 网络

使用Python实现图像LBP特征提取的操作方法

《使用Python实现图像LBP特征提取的操作方法》LBP特征叫做局部二值模式,常用于纹理特征提取,并在纹理分类中具有较强的区分能力,本文给大家介绍了如何使用Python实现图像LBP特征提取的操作方... 目录一、LBP特征介绍二、LBP特征描述三、一些改进版本的LBP1.圆形LBP算子2.旋转不变的LB

Redis消息队列实现异步秒杀功能

《Redis消息队列实现异步秒杀功能》在高并发场景下,为了提高秒杀业务的性能,可将部分工作交给Redis处理,并通过异步方式执行,Redis提供了多种数据结构来实现消息队列,总结三种,本文详细介绍Re... 目录1 Redis消息队列1.1 List 结构1.2 Pub/Sub 模式1.3 Stream 结

C# Where 泛型约束的实现

《C#Where泛型约束的实现》本文主要介绍了C#Where泛型约束的实现,文中通过示例代码介绍的非常详细,对大家的学习或者工作具有一定的参考学习价值,需要的朋友们下面随着小编来一起学习学习吧... 目录使用的对象约束分类where T : structwhere T : classwhere T : ne

将Java程序打包成EXE文件的实现方式

《将Java程序打包成EXE文件的实现方式》:本文主要介绍将Java程序打包成EXE文件的实现方式,具有很好的参考价值,希望对大家有所帮助,如有错误或未考虑完全的地方,望不吝赐教... 目录如何将Java程序编程打包成EXE文件1.准备Java程序2.生成JAR包3.选择并安装打包工具4.配置Launch4

MySQL索引的优化之LIKE模糊查询功能实现

《MySQL索引的优化之LIKE模糊查询功能实现》:本文主要介绍MySQL索引的优化之LIKE模糊查询功能实现,本文通过示例代码给大家介绍的非常详细,感兴趣的朋友一起看看吧... 目录一、前缀匹配优化二、后缀匹配优化三、中间匹配优化四、覆盖索引优化五、减少查询范围六、避免通配符开头七、使用外部搜索引擎八、分

Python实现特殊字符判断并去掉非字母和数字的特殊字符

《Python实现特殊字符判断并去掉非字母和数字的特殊字符》在Python中,可以通过多种方法来判断字符串中是否包含非字母、数字的特殊字符,并将这些特殊字符去掉,本文为大家整理了一些常用的,希望对大家... 目录1. 使用正则表达式判断字符串中是否包含特殊字符去掉字符串中的特殊字符2. 使用 str.isa

Spring Boot 集成 Quartz并使用Cron 表达式实现定时任务

《SpringBoot集成Quartz并使用Cron表达式实现定时任务》本篇文章介绍了如何在SpringBoot中集成Quartz进行定时任务调度,并通过Cron表达式控制任务... 目录前言1. 添加 Quartz 依赖2. 创建 Quartz 任务3. 配置 Quartz 任务调度4. 启动 Sprin

Android实现悬浮按钮功能

《Android实现悬浮按钮功能》在很多场景中,我们希望在应用或系统任意界面上都能看到一个小的“悬浮按钮”(FloatingButton),用来快速启动工具、展示未读信息或快捷操作,所以本文给大家介绍... 目录一、项目概述二、相关技术知识三、实现思路四、整合代码4.1 Java 代码(MainActivi