PyTorch深度学习实战 |手算GCN (图神经网络)模型

简介: 本文介绍了使用PyTorch实现图神经网络(GNN)处理分子结构数据的实战方法。主要内容包括:1) GNN的基本原理,通过节点特征矩阵和邻接矩阵处理图结构数据;2) 分子图的表示方式,将SMILES字符串转换为PyTorch Geometric图对象;3) 图卷积运算过程,包括特征变换和邻接特征聚合;4) 代码实现示例,构建包含GCN层和全局池化的模型,对乙醇分子进行特征提取和分类预测。文章通过具体案例展示了GNN在化学领域的应用,为读者提供了从理论到实践的完整指导。

 💻图神经网络

      图神经网络(Graph Neural Network,GNN)是一种专门用于处理图结构数据的深度学习方

法。与传统的神经网络主要处理规则结构的数据(如图像和文本)不同,GNN能够处理各种不规

则的数据结构,如社交网络、分子结构等。GNN通过在图上定义节点之间的连接关系,利用节点

的邻居信息来更新节点的表示,实现对整个图的信息传递和学习。

image.gif


📘分子的图结构

图神经网络(GNN)处理的输入是图(Graph),而不是传统的像素矩阵或序列。因此,第一步

是将我们的目标分子——乙醇,抽象地转化为一个数学图。

节点与邻接关系

为了简化手算过程,我们只关注重原子:两个碳原子和一个氧原子。我们将氢原子的影响体现在

节点的特征中。

邻接矩阵 : 描述节点之间的连接关系(化学键)。

image.gif

节点特征矩阵:是 GCN 的输入。为了演示目的,我们为每个原子分配一个简单的特

征向量(例如一个独热编码和它的连接度):

image.gif

🔬 真实项目中分子图的表示方式

在真实的分子图神经网络项目中,虽然图的基本原理(节点和边)是一样的,但节点特征图的复

杂度会有巨大的差异。在 PyTorch Geometric (PyG) 或 Deep Graph Library (DGL) 等专业 GNN 库

中处理分子(例如基于 RDKit 的分子表示)时,图的定义会丰富得多,最后我们会详细的介绍一

下这部分的内容。


📘图卷积

    图卷积的输入数据是节点特征矩阵H和邻接矩阵A,下图我们展示了3个节点的图,每个节点特

征数为3的特征,图卷积在计算的时候有两个关键的步骤,分别是节点特征的线性变换邻接特征

的聚合

GCN 单层的核心公式是:

节点特征的线性变换

            就是使用线性层,对节点特征矩阵进行线性变换,提取特征,节点的特征数目发生变化。

X' 现在代表了每个原子经过权重变换后的新特征。计算过程如下:

image.gif

邻接特征的聚合

    这里为了方便理解,邻接矩阵没有使用稀疏表示,就是使用邻接矩阵adj进行特征聚合,将相邻

节点中的特征信息,传导到该节点上。

新的特征就等于聚合邻居信息变换特征的乘积

聚合邻居信息

这个矩阵中的每一行和每一列的元素,定义了消息如何从邻居节点传递并加权求和

image.gif

总结一下:一次的图卷积,本质上是特征的变换(或者说是特征数的变化),在这个过程中,节点

特征矩阵数据,包含每个节点的属性信息,每层都会通过 GCN 运算更新。归一化邻接矩阵结构

定义图的拓扑结构和消息传递路径,它是固定的。


📘分子图的表达方式

🔬 真实项目中分子图的表示方式

在真实项目中,分子表示的流程是:

SMILES 字符串      RDKit 解析     PyG/DGL 图对象

安装相关的包:

pip install rdkit

image.gif

pip install torch_geometric

image.gif

# 假设 PyTorch, RDKit, PyTorch Geometric (PyG) 库已安装
import torch
from rdkit import Chem
from torch_geometric.utils.smiles import from_smiles  # PyG中用于SMILES转换的实用函数 (或使用更早版本的'from_rdkit')
# --- 1. 定义和转换 ---
smiles_ethanol = "CCO" 
# 使用 PyG 的封装函数,一步完成解析、特征提取和图结构构建
# 这个函数内部自动完成了原子特征编码、键索引构建、以及 PyTorch 张量转换。
ethanol_data = from_smiles(smiles_ethanol)
# --- 2. 展示结果 ---
print("=" * 40)
print(f"乙醇分子 SMILES: {smiles_ethanol}")
print("PyG Data 对象结构 (封装结果)")
print("=" * 40)
# PyG Data 对象概览
print(ethanol_data) 
print("\n--- 关键张量尺寸分析 ---")
print(f"节点特征矩阵 (x): {ethanol_data.x.shape}")
print(f"邻接信息 (edge_index): {ethanol_data.edge_index.shape}")
print("-" * 40)

image.gif

========================================

乙醇分子 SMILES: CCO

PyG Data 对象结构 (封装结果)

========================================

Data(x=[3, 9], edge_index=[2, 4], edge_attr=[4, 3], smiles='CCO')


--- 关键张量尺寸分析 ---

节点特征矩阵 (x): torch.Size([3, 9])

邻接信息 (edge_index): torch.Size([2, 4])

----------------------------------------

📈 PyG Data 对象结构解读

3 (行): 代表图中的节点数量, 模型忽略了 6 个氢原子。9 (列): 代表每个节点的特征维度特征提取

器为每个重原子编码了 9 种不同的化学属性(例如:原子类型、价态、电荷、隐式氢数等)。2

(行): PyG 的固定格式,第一行是源节点索引,第二行是目标节点索引。4 (列): 代表有向边的总

数。


图卷积神经网络

以乙醇分子为例,模拟前向传播

原始3×9大小的张量

经过隐藏层通道数为16的图卷积,得到3×16大小的张量

再经过隐藏层通道数为16的图卷积,得到3×16大小的张量

再经过一个全局平均池化,得到1×16大小的特征矩阵,全局平均将节点特征聚合为图特征(每个图一个数来表示)

然后再通过一个线性层变成1×16大小的张量

🔍 代码实现

import torch
import torch.nn.functional as F
from torch_geometric.nn import GCNConv, global_mean_pool
from torch_geometric.data import Data
import sys
# 尝试导入 RDKit 和 PyG 转换工具
from rdkit import Chem
from torch_geometric.utils.smiles import from_smiles
# --- 1. 真实数据准备 (使用 PyG 封装函数) ---
smiles_ethanol = "CCO" 
# 一步转换:生成包含所有 9 个原子(C, C, O, 6H)的图结构
ethanol_data = from_smiles(smiles_ethanol)
 # 修复:确保节点特征是浮点类型 (解决 RuntimeError)
ethanol_data.x = ethanol_data.x.float()
# 生成 Batch Tensor:9 个节点都属于同一个图 (batch size=1)
N_NODES = ethanol_data.x.shape[0]
ethanol_data.batch = torch.zeros(N_NODES, dtype=torch.long)
# --- 2. 展示输入数据结构 ---
print("=" * 60)
print(f"乙醇分子 SMILES: {smiles_ethanol}")
print("PyG Data 对象结构 (GCN 模型输入 - 真实配置)")
print("=" * 60)
print(f"节点数 N: {ethanol_data.x.shape[0]}")
print(f"节点特征矩阵 (x): {ethanol_data.x.shape}")
print(f"邻接信息 (edge_index): {ethanol_data.edge_index.shape}")
print("-" * 60)
# --- 3. 定义 GCN 模型类 (SimpleGCN) ---
class SimpleGCN(torch.nn.Module):
    def __init__(self, num_node_features, hidden_channels, num_classes):
        super().__init__()
        self.conv1 = GCNConv(num_node_features, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, hidden_channels)
        self.lin = torch.nn.Linear(hidden_channels, num_classes)
    def forward(self, x, edge_index, batch):
        print(f"\n[A] 初始输入 x (H(0)): {x.shape}")
        # GCN 层 1
        x = self.conv1(x, edge_index)
        print(f"[B] GCNConv 1 输出 (H(1)): {x.shape}")
        x = F.relu(x)
        # GCN 层 2
        x = self.conv2(x, edge_index)
        print(f"[C] GCNConv 2 输出 (H(2)): {x.shape}")
        x = F.relu(x)
        # 全局读出/池化层
        x = global_mean_pool(x, batch)
        print(f"[D] Global Mean Pool 输出: {x.shape} <--- **图级特征**")
        # 线性分类层
        x = self.lin(x)
        print(f"[E] 最终分类层输出: {x.shape} <--- **预测结果**")
        return x
# --- 4. 模型实例化与运行 ---
# 定义模型参数
INPUT_DIM = ethanol_data.x.shape[1] # 自动获取真实/模拟的特征维度 (例如:11)
HIDDEN_DIM = 16 
OUTPUT_DIM = 1 
# 实例化模型
model = SimpleGCN(
    num_node_features=INPUT_DIM, 
    hidden_channels=HIDDEN_DIM, 
    num_classes=OUTPUT_DIM
)
# 执行前向传播
print("=" * 60)
print(f"【Simple GCN 前向传播过程(隐藏层维度 D_hidden={HIDDEN_DIM})】")
print("=" * 60)
output = model(
    ethanol_data.x, 
    ethanol_data.edge_index, 
    ethanol_data.batch
)

image.gif

============================================================

乙醇分子 SMILES: CCO

PyG Data 对象结构 (GCN 模型输入 - 真实配置)

============================================================

节点数 N: 3

节点特征矩阵 (x): torch.Size([3, 9])

邻接信息 (edge_index): torch.Size([2, 4])

------------------------------------------------------------

============================================================

【Simple GCN 前向传播过程(隐藏层维度 D_hidden=16)】

============================================================


[A] 初始输入 x (H(0)): torch.Size([3, 9])

[B] GCNConv 1 输出 (H(1)): torch.Size([3, 16])

[C] GCNConv 2 输出 (H(2)): torch.Size([3, 16])

[D] Global Mean Pool 输出: torch.Size([1, 16]) <--- **图级特征**

[E] 最终分类层输出: torch.Size([1, 1]) <--- **预测结果**


目录
相关文章
|
3月前
|
机器学习/深度学习 编解码 算法
PyTorch深度学习实战 |手算​​U-net
本文详细解析了U-Net网络架构及其在医学图像分割中的应用。重点对比了U-Net与FCN的核心区别:U-Net采用特征拼接(Concat)保留所有层级信息,而FCN使用特征相加(Add)进行融合。文章深入剖析了U-Net的编码器-瓶颈-解码器结构,解释了其独特的裁剪拼接机制和Overlap-tile策略,并提供了完整的PyTorch实现代码。现代U-Net通过SamePadding实现了输入输出尺寸一致,显著提升了分割精度。文章还探讨了弹性形变数据增强和带空间权重的损失函数设计,为医学图像分析提供了实用解决
317 2
|
3月前
|
存储 人工智能 自然语言处理
拒绝“大模型幻觉”:一文彻底搞懂 RAG(检索增强生成)技术全流程
本文深入解析RAG(检索增强生成)技术,直击大模型落地私有知识场景的核心痛点——如何让LLM精准、低成本、高时效地基于企业文档作答。从文本分片、向量化索引,到召回重排、增强生成,系统拆解五大关键步骤,揭示RAG作为“AI外挂”的底层逻辑与工程实践精髓。
拒绝“大模型幻觉”:一文彻底搞懂 RAG(检索增强生成)技术全流程
|
3月前
|
弹性计算 前端开发 Ubuntu
阿里云服务器ECS的租用教程和简单的前端页面部署
本文详解阿里云学生福利领取(含300元卡券)及ECS轻量服务器选购与部署全流程:涵盖学生机免费申领、配置选型建议(Ubuntu/CentOS/Windows)、安全组设置、Nginx安装、网页部署及Xshell远程连接等实操步骤,新手友好。
450 8
|
3月前
|
人工智能 开发工具 开发者
学习AI Agent编程-第一天-MCP基础
本文精炼解析MCP(Model Context Protocol):它不是新模型,而是让AI Agent运行时动态增删工具的协议。通过MCP Server(工具实现)、Client(SDK封装)与Host(Agent应用)三组件协作,解决传统`bind_tools`静态绑定的局限。附完整可运行示例,助你快速掌握80%核心用法。(239字)
570 1
|
3月前
|
机器学习/深度学习 自然语言处理 PyTorch
PyTorch深度学习实战 |手动计算 Transformer和完整的代码实现
本文介绍了基于PyTorch实现Transformer模型的完整过程。主要内容包括:1)Transformer架构的核心组件实现,如多头注意力机制、位置前馈网络、位置编码等;2)模型构建步骤,包括词嵌入层、编码器/解码器块和输出层的实现;3)完整的训练流程,包含数据处理、损失计算和参数优化;4)评估方法验证模型性能。文章通过代码示例详细展示了如何从零开始构建Transformer,并应用于机器翻译任务,同时对模型各层的输入输出维度进行了说明。该实现可作为深度学习实践者学习Transformer架构的实用指南
306 0
|
3月前
|
机器学习/深度学习 数据采集 人工智能
田间杂草检测数据集分享(适用于YOLO系列深度学习分类检测任务)
本数据集含4000张真实农田图像(小麦/玉米/水稻田),YOLO格式标注杂草目标,覆盖多天气、光照与视角,适用于YOLO系列等目标检测模型训练,助力智能除草与精准农业研究。(239字)
498 16
|
3月前
|
人工智能 机器人 芯片
人工智能|YOLOv8实战
本内容为安全帽检测实战项目,基于YOLOv8模型,涵盖Kaggle数据获取、自定义yaml配置、模型训练(yolo_train.py)与测试(yolo_test.py),并提供服务器(FastAPI+Docker)、边缘(Jetson+TensorRT)及国产嵌入式(RK3588+RKNN)三类部署方案,支持工业场景实时智能识别。(239字)
543 1
|
1月前
|
人工智能 架构师 安全
【AI时代软件项目管理系列】2. 从工具到数字员工:AI 在项目团队中的角色演进
AI 正从个人助手演进为能够规划、调用工具并完成任务的 Agent,进一步形成多 Agent 协作与软件工厂。角色变化带来的不只是效率提升,更是团队结构、任务分配、进度衡量和质量管理方式的重构。项目经理需要把模型、Agent、知识库和自动化流程纳入统一管理,通过明确任务契约、权限边界、质量门禁与人工验收,确保 AI 执行可控、结果可验证、责任可追溯。
275 0
|
2月前
|
存储 API 数据处理
阿里云智能媒体管理(IMM)对接使用全攻略:从开通到生产级实践
本文全面解析阿里云智能媒体管理(IMM)的对接与使用。首先介绍IMM的产品定位与服务架构,阐明其与OSS的深度集成关系。然后详细说明开通服务、创建项目、绑定Bucket的全流程操作,并给出Java和Python两种主流语言的SDK初始化与API调用示例。接着深入讲解文档格式转换、文档预览、视频截帧、图片智能检测等核心功能的实现方法,涵盖同步与异步处理两种模式。同时针对权限配置、新旧版本差异、计费规则、性能优化等关键问题进行专项剖析,帮助读者构建生产级的媒体处理能力。全文基于新版IMM(API版本2020-09-30)撰写,适合开发者、架构师及技术决策者阅读。
|
2月前
|
域名解析 负载均衡 网络协议
阿里云DNS云解析:公网权威解析个人版费用19.9元1年,支持功能、安全配置及续费说明
阿里云DNS个人版年费19.9元(原价48元),限个人开发者,支持DNSSEC、智能解析、URL转发等,全球百余节点,100%可用性保障;续费按原价48元/年,不自动延续优惠。阿里云云解析DNS官网:https://t.aliyun.com/U/h4bNRD