目录

mGATGNN

misaraty 更新 | 2026-10-06
前言
下载:mGATGNN。

mGATGNN

mGATGNN 是一个基于 PyTorch 和 PyTorch Geometric 的独立 GATGNN 实现,用于根据 CIF 晶体结构预测材料带隙。

代码保留了 GATGNN 的主要组成部分,包括周期性晶体图构建、高斯距离展开、增强型多头图注意力(AGAT)、组成引导的全局注意力、图级池化和回归头。代码使用 pymatgen 解析 CIF 并构建周期近邻,不依赖原始 GATGNN 仓库。

原始 GATGNN 数据流程采用的92维 CGCNN 元素描述符已经直接嵌入脚本,因此不需要额外提供 atom_init.json,也不需要安装或下载 CGCNN 仓库。

模型结构

默认模型为每个原子选取12个周期近邻,使用41维高斯距离特征、3层 AGAT、4个注意力头、64维隐藏特征,并使用103维元素组成向量计算全局注意力。常用的晶体图参数和模型参数均可在脚本开头修改。

训练采用 MSE 损失、AdamW、目标标准化、验证集 RMSE 模型选择、ReduceLROnPlateau、早停、梯度裁剪,并在支持的 CUDA 设备上使用 BF16 自动混合精度。脚本还保留了可选的 Optuna 超参数优化开关。

环境依赖

1
pip install torch torch-geometric pymatgen numpy pandas openpyxl scikit-learn matplotlib tqdm

仅当 USE_OPTUNA = True 时需要安装 Optuna:

1
pip install optuna

数据准备

在脚本所在目录中准备以下文件:

1
2
3
4
5
6
mGATGNN_v1.py
data.xlsx
cif/
|-- 1.cif
|-- 2.cif
|-- 3.cif

脚本读取 data.xlsx 的前两列:

cif bandgap
1 1.23
2 0.87
3 2.15

第一列中的 1 对应 ./cif/1.cif。如果第一列没有包含 .cif 后缀,脚本会自动补充。

使用方法

1
python mGATGNN_v1.py

脚本会自动完成 CIF 检查、晶体图缓存、固定的80/10/10训练集/验证集/测试集划分、模型训练、早停、评估和绘图。当 CUDA 可用时自动使用 GPU,否则使用 CPU。

所有结果保存在由 MODEL_NAME 和 RUN_VERSION 确定的目录中。按照默认设置,输出目录为:

1
2
3
4
5
6
7
8
GATGNN_v1/
|-- GATGNN_best.pt
|-- figure/
|-- dat/
|-- table/
|-- log/
|-- split/
|-- cache/

输出内容包括训练集、验证集和测试集的 MAE、RMSE 与 R2,各数据集及合并奇偶图和对应数据,RMSE 训练曲线及其数据,固定数据划分,最佳模型权重,缓存晶体图和完整训练日志。

保存的 checkpoint 包含模型权重、模型与晶体图配置、目标标准化参数、嵌入式原子描述符、最佳 epoch、最佳验证集 RMSE、目标信息和训练参数。脚本还提供 load_trained_model() 和 predict_cifs(),用于预测新的 CIF 晶体结构。

引用

原始 GATGNN 文献:

本工作:

论文正式发表后补充。