mCrystalFramer
mCrystalFramer
mCrystalFramer 是一个基于纯 PyTorch 的独立 CrystalFramer 风格实现,用于根据 CIF 晶体结构预测材料带隙。
代码保留了 CrystalFramer 的主要思想,包括周期镜像注意力、距离相关注意力偏置、高斯距离编码、各注意力头独立的动态 max frame、基于 frame 的角度编码、Transformer 风格残差模块、晶体级平均池化和标量回归头。代码使用 pymatgen 解析 CIF,不依赖 PyTorch Geometric、DGL、JARVIS、CuPy、自定义 CUDA kernel 或原始 CrystalFramer 仓库。
这个独立版本使用原生 PyTorch 张量运算实现有限周期镜像求和。CUDA 可用时仍会自动使用 GPU,但其数值结果和运行速度可能与原始融合 CuPy/CUDA 实现存在差异。
环境依赖
|
|
仅当 USE_OPTUNA = True 时需要安装 Optuna:
|
|
数据准备
在脚本所在目录中准备以下文件:
|
|
脚本读取 data.xlsx 的前两列:
| cif | bandgap |
|---|---|
| 1 | 1.23 |
| 2 | 0.87 |
| 3 | 2.15 |
第一列中的 1 对应 ./cif/1.cif。脚本可以识别整数、1.0 这类整数值浮点数、不带 .cif 后缀的名称以及完整 CIF 文件名。
使用方法
|
|
脚本会自动检查 CIF 文件,使用 SEED = 42 完成固定的 80/10/10 训练集/验证集/测试集划分,仅根据训练集统计量标准化目标值,使用 MSELoss 训练,根据验证集 RMSE 调整学习率和执行早停,重新加载最佳 checkpoint,评估三个数据集并生成图表与数据文件。当 CUDA 可用时自动使用 GPU 和 BF16 自动混合精度,否则使用 CPU。
所有结果保存在由 MODEL_NAME 和 RUN_VERSION 确定的目录中。按照默认设置,输出目录为:
|
|
输出内容包括训练集、验证集和测试集的 MAE、RMSE 与 R2,各数据集及合并奇偶图与对应数据,训练集/验证集 RMSE 曲线及对应数据,可复用的固定数据划分,最佳模型 checkpoint 和完整训练日志。Checkpoint 中保存模型配置、目标标准化参数、最佳 epoch、最佳验证集 RMSE 和图构建配置。
脚本还提供 load_trained_model() 和 predict_cifs(),可供其他 Python 程序加载已保存的 checkpoint 并预测新的 CIF 结构。
显存说明
对于原子数较多的结构,周期注意力和基于 frame 的角度编码可能占用较多 GPU 显存。如果出现 CUDA 显存不足,建议按照以下顺序降低参数:
|
|
建议尽可能保留 LATTICE_RANGE = 1,因为设置为零会删除非中心周期镜像。ANGLE_BASIS_DIM 必须能够被 3 整除,MODEL_DIM 必须能够被 NUM_HEADS 整除。
引用
原始 CrystalFramer 文献:
本工作:
论文正式发表后补充。