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 超参数优化开关。
环境依赖
|
|
仅当 USE_OPTUNA = True 时需要安装 Optuna:
|
|
数据准备
在脚本所在目录中准备以下文件:
|
|
脚本读取 data.xlsx 的前两列:
| cif | bandgap |
|---|---|
| 1 | 1.23 |
| 2 | 0.87 |
| 3 | 2.15 |
第一列中的 1 对应 ./cif/1.cif。如果第一列没有包含 .cif 后缀,脚本会自动补充。
使用方法
|
|
脚本会自动完成 CIF 检查、晶体图缓存、固定的80/10/10训练集/验证集/测试集划分、模型训练、早停、评估和绘图。当 CUDA 可用时自动使用 GPU,否则使用 CPU。
所有结果保存在由 MODEL_NAME 和 RUN_VERSION 确定的目录中。按照默认设置,输出目录为:
|
|
输出内容包括训练集、验证集和测试集的 MAE、RMSE 与 R2,各数据集及合并奇偶图和对应数据,RMSE 训练曲线及其数据,固定数据划分,最佳模型权重,缓存晶体图和完整训练日志。
保存的 checkpoint 包含模型权重、模型与晶体图配置、目标标准化参数、嵌入式原子描述符、最佳 epoch、最佳验证集 RMSE、目标信息和训练参数。脚本还提供 load_trained_model() 和 predict_cifs(),用于预测新的 CIF 晶体结构。
引用
原始 GATGNN 文献:
本工作:
论文正式发表后补充。