【MIIT program】Feature add SGEquiDiff - #287
Closed
learncat163 wants to merge 321 commits into
Closed
Conversation
add atom type diffusion
* chore: delete duplicate code * chore: change save frequency
* feat: change diffusion code for mattergen * fix: scale shared grads only works on gemnet
* chore: change save path in config * chore: annotate useless code * update config * add guiduce for mattergen
* chore: annotate useless code * [feat]: add reconstruct metrics
add diffcsp++ with types and fixe the sample function of diffcsp++
* feat: add dimenet model * feat: add radius_graph_pbc_wrapper for generation metirc * chore: rename file * feat: add property metric for generation
* chore: update config * fix: update n_sample values * fix: atom start index * feat: add eval function
* Delete electronic_structure/configs/infgcn/infgcn_omol25_trimmed.yaml * Update README.md * Update infgcn_omol25_MC_5k_trimmed.yaml * Delete electronic_structure/configs/infgcn/infgcn_omol25.yaml
* Revise copyright and authorship statements Updated copyright information and authorship details. * Revise copyright and ownership details in mattergen.py Updated copyright notice and ownership information. * Update property_embeddings.py * Update copyright and ownership information in globals.py
…e#274) Co-authored-by: leeleolay <leeleolay@example.com>
* Update docs: rename MetaX logo and add new partner logos * update readme --------- Co-authored-by: leeleolay <leeleolay@example.com>
|
Thanks for your contribution! |
Author
3.3 采样指标对齐采样指标对接结果
采样的对齐步骤比较复杂,主要分离为2步:
3.3.1 PT的晶体采样PT晶体采样Python脚本,点击后展开脚本#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
PyTorch 原始环境晶体生成脚本。
在原始 sgequidiff 的venv 环境中运行,使用原始 PyTorch 代码生成晶体样本。
使用方法:
source ~/opensource/cailiao/sgequidiff/.venv/bin/activate
python generate_samples_pt.py --num_samples 64 --output_dir ../outputs/pt_samples
"""
import argparse
import json
import os
import sys
import time
from pathlib import Path
def parse_args():
parser = argparse.ArgumentParser(description="PT 晶体生成")
parser.add_argument("--num_samples", type=int, default=64, help="生成样本数")
parser.add_argument("--batch_size", type=int, default=32, help="批大小")
parser.add_argument("--temperature", type=float, default=1.0, help="采样温度")
parser.add_argument("--seed", type=int, default=42, help="随机种子")
parser.add_argument("--ckpt_dir", type=str, required=True, help="PT 权重目录")
parser.add_argument("--output_dir", type=str, required=True, help="输出目录")
parser.add_argument("--dataset", type=str, default="mp_20", choices=["mp_20", "mpts_52"])
return parser.parse_args()
def main():
args = parse_args()
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
# 记录日志
log_path = output_dir / "generation.log"
sys.stdout = open(log_path, "w")
sys.stderr = sys.stdout
print(f"[{time.strftime('%Y-%m-%d %H:%M:%S')}] PT 晶体生成开始")
print(f" 参数: num_samples={args.num_samples}, batch_size={args.batch_size}, "
f"temperature={args.temperature}, seed={args.seed}")
print(f" 权重: {args.ckpt_dir}")
print(f" 输出: {args.output_dir}")
# 添加原始项目路径(从脚本位置上溯到项目根)
raw_project = Path(__file__).resolve().parents[4] / "sgequidiff-raw"
sys.path.insert(0, str(raw_project))
os.chdir(str(raw_project))
# 导入原始模块
from scripts.generate_crystals import main as raw_main
# 构造 raw_main 的参数
sys.argv = [
"generate_crystals.py",
"--num_samples", str(args.num_samples),
"--batch_size", str(args.batch_size),
"--ckpt_dir", args.ckpt_dir,
"--load_best_submodules",
"--temperature", str(args.temperature),
"--seed", str(args.seed),
"--save_method", "cif",
"--save_dir", str(output_dir / "cifs"),
]
if args.dataset == "mpts_52":
sys.argv += ["--dataset_name", "mpts_52"]
print(f"\n启动原始生成脚本...")
raw_main()
# 保存元数据
meta = {
"framework": "pt",
"dataset": args.dataset,
"num_samples": args.num_samples,
"temperature": args.temperature,
"seed": args.seed,
"ckpt_dir": args.ckpt_dir,
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
}
with open(output_dir / "metadata.json", "w") as f:
json.dump(meta, f, indent=2)
print(f"\n[{time.strftime('%Y-%m-%d %H:%M:%S')}] PT 晶体生成完成")
print(f" 输出目录: {args.output_dir}")
if __name__ == "__main__":
main()3.3.2 PD的晶体采样PD晶体采样Python脚本,点击后展开脚本#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
Paddle 环境晶体生成脚本(独立进程隔离版)。
每生成一个晶体,都启动一个独立子进程,进程退出后操作系统自动回收全部 GPU 显存。 不然会爆显存 OOM!!!!
主进程本身不加载模型、不分配 GPU 资源,仅负责调度和日志记录。
架构:
main() -- 主进程: 调度循环,不碰 GPU
|
+-- subprocess python generate_samples_pd.py --worker 0
+-- subprocess python generate_samples_pd.py --worker 1
+-- ...
每个子进程: 加载模型 -> 采样 -> 保存 CIF -> 进程退出 -> GPU 显存彻底释放
使用方法:
conda activate ppmat
python generate_samples_pd.py --num_samples 64 --output_dir ./output
"""
import argparse
import json
import os
import subprocess
import sys
import time
from pathlib import Path
def parse_args():
parser = argparse.ArgumentParser(
description="PD 晶体生成(独立进程隔离,每晶体独立子进程)")
parser.add_argument("--num_samples", type=int, default=16,
help="生成样本数 (默认 16)")
parser.add_argument("--temperature", type=float, default=1.0,
help="采样温度 (默认 1.0)")
parser.add_argument("--seed", type=int, default=42,
help="随机种子 (默认 42)")
parser.add_argument("--output_dir", type=str, required=True,
help="输出目录")
parser.add_argument("--dataset", type=str, default="mp_20",
choices=["mp_20", "mpts_52"],
help="数据集 (默认 mp_20)")
parser.add_argument("--diffusion_snr", type=float, default=0.4,
help="扩散 SNR (默认 0.4)")
parser.add_argument("--max_step_size", type=float, default=1e6,
help="扩散采样最大步长")
# --- worker 模式(由主进程自动调用,用户无需手动指定) ---
parser.add_argument("--worker", type=int, default=None,
help="[内部] worker 模式: 生成第 N 个晶体后退出")
return parser.parse_args()
# ---------------------------------------------------------------------------
# Worker: 在独立子进程中运行,生成单个晶体后退出,GPU 显存随之彻底释放
# ---------------------------------------------------------------------------
def worker_main(args):
"""子进程入口: 构建模型 -> 采样 -> 保存 CIF -> 退出。"""
project_root = str(Path(__file__).resolve().parents[4])
if project_root not in sys.path:
sys.path.insert(0, project_root)
import paddle
from ppmat.models.sgequidiff.weight_utils import load_pretrained_weights
from ppmat.models.sgequidiff.crystal_sampler import (
CrystalSampler, CrystalSamplerConfig,
)
from ppmat.models.sgequidiff.diffusion_model import (
EquivariantDiffusionModelConfig,
)
from ppmat.models.sgequidiff.lattice_sampler import LatticeSamplerConfig
from ppmat.models.sgequidiff.wyckoff_transformer import (
WyckoffElementTransformerConfig,
)
from ppmat.models.sgequidiff.non_equivariant_drift_modules import GNNConfig
from ppmat.models.sgequidiff.embedding_utils import (
set_global_embedding_tools,
)
from ppmat.models.sgequidiff.constants import (
lattice_parameter_ranges, chemical_symbols,
)
idx = args.worker
# 设置随机种子(每个 worker 不同)
paddle.seed(args.seed + idx)
# 初始化全局 embedding 工具
set_global_embedding_tools()
# 构建模型配置
lr = lattice_parameter_ranges.get(
args.dataset, lattice_parameter_ranges["mp_20"],
)
gnn_cfg = GNNConfig(
num_plane_wave_freqs=96,
num_cartesian_distance_gaussians=96,
edge_hidden_dim=128,
atom_hidden_dim=256,
use_vpa=True,
use_graph_norm=True,
num_msg_pass_steps=5,
cutoff=10.0,
use_frac_coords_in_node_emb=True,
dataset_name=args.dataset,
)
diff_cfg = EquivariantDiffusionModelConfig(
model_type="gnn",
num_timesteps=1000,
noise_scheduler_num_monte_carlo_samples=2500,
num_wn_lattice_translations=3,
sigma_min=0.002,
sigma_max=0.5,
time_emb_dim=128,
num_plane_wave_freqs=96,
gnn_config=gnn_cfg,
)
lattice_cfg = LatticeSamplerConfig(
input_dimension=128,
hidden_dimension=256,
min_lattice_length=lr["min_lattice_length"],
max_lattice_length=lr["max_lattice_length"],
min_lattice_angle=lr["min_lattice_angle"],
max_lattice_angle=lr["max_lattice_angle"],
)
we_cfg = WyckoffElementTransformerConfig(
hidden_dim=256,
dataset_name=args.dataset,
num_heads=4,
num_hidden_layers=1,
dropout_rate=0.0,
)
sampler_cfg = CrystalSamplerConfig(
diffusion_model_config=diff_cfg,
lattice_model_config=lattice_cfg,
transformer_config=we_cfg,
)
# 构建模型 & 加载权重
model = CrystalSampler(sampler_cfg)
model.eval()
load_pretrained_weights(model, dataset_name=args.dataset, verbose=False)
# 采样
try:
crystals = model.sample_crystal(
batch_size=1,
diffusion_snr=args.diffusion_snr,
temperature=args.temperature,
)
except Exception as e:
print(f"FAIL:{idx}:sample_error:{e}", flush=True)
return
if not crystals or len(crystals) == 0:
print(f"FAIL:{idx}:empty_result", flush=True)
return
# 保存 CIF
c = crystals[0]
lattice = (
c.conventional_lattice_lengths.tolist()
+ c.conventional_lattice_angles.tolist()
)
elements = [chemical_symbols[i] for i in c.element_indices.tolist()]
coords = c.conventional_frac_coords.tolist()
cif_path = Path(args.output_dir) / "cifs" / f"gen_{idx:05d}.cif"
_save_cif(cif_path, lattice, elements, coords, c.space_group_number)
# 输出统计信息(主进程可捕获)
real_atoms = sum(1 for e in elements if e not in ("X", "X0+", ""))
print(f"OK:{idx}:atoms={real_atoms}", flush=True)
# 显式清理(虽然进程即将退出,但保险起见)
del model, crystals
paddle.device.cuda.empty_cache()
# ---------------------------------------------------------------------------
# CIF 保存(混合占位模式,与 PT 版本对齐)
# ---------------------------------------------------------------------------
def _save_cif(path, lattice_params, elements, frac_coords, space_group):
"""将晶体信息保存为 CIF 文件。
相同分数坐标的原子合并为混合占位 (mixed-occupancy) 位点,
使 CIF 的位点数与 PT 版本一致,消除评估偏差。
"""
from collections import Counter, defaultdict
from pymatgen.core import Element, Lattice, Structure
# 过滤占位元素
valid_pairs = []
for i, el_name in enumerate(elements):
if el_name in ("X", "X0+", "") or len(el_name) == 0:
continue
try:
el = Element(el_name)
coord = tuple(frac_coords[i])
valid_pairs.append((el, coord))
except Exception:
continue
if not valid_pairs:
return
# 按坐标分组,合并相同坐标的不同元素为混合占位
site_map = defaultdict(list)
for el, coord in valid_pairs:
site_map[coord].append(el)
group_species = []
group_coords = []
for coord, elements_list in site_map.items():
el_counts = Counter(elements_list)
total = sum(el_counts.values())
if total == 1:
group_species.append(elements_list[0])
else:
group_species.append(
{el: count / total for el, count in el_counts.items()}
)
group_coords.append(list(coord))
a, b, c = lattice_params[0], lattice_params[1], lattice_params[2]
alpha, beta, gamma = (
lattice_params[3], lattice_params[4], lattice_params[5],
)
lattice = Lattice.from_parameters(a, b, c, alpha, beta, gamma)
structure = Structure(
lattice, group_species, group_coords, coords_are_cartesian=False,
)
from pymatgen.io.cif import CifWriter
writer = CifWriter(structure, symprec=None)
path.parent.mkdir(parents=True, exist_ok=True)
writer.write_file(str(path))
# ---------------------------------------------------------------------------
# 主进程: 调度循环,不加载 GPU 资源
# ---------------------------------------------------------------------------
def scheduler_main(args):
"""主进程: 逐个启动子进程,每个子进程生成一个晶体后退出。"""
output_dir = Path(args.output_dir)
cifs_dir = output_dir / "cifs"
cifs_dir.mkdir(parents=True, exist_ok=True)
log_path = output_dir / "generation.log"
with open(log_path, "w") as log_file:
def log(msg):
ts = time.strftime("%Y-%m-%d %H:%M:%S")
line = f"[{ts}] {msg}"
print(line, flush=True)
log_file.write(line + "\n")
log_file.flush()
log("=" * 60)
log("PD 晶体生成(独立进程隔离模式)")
log("=" * 60)
log(f" num_samples = {args.num_samples}")
log(f" dataset = {args.dataset}")
log(f" temperature = {args.temperature}")
log(f" seed = {args.seed}")
log(f" output_dir = {args.output_dir}")
log(f" 每个晶体使用独立子进程,进程退出即释放全部 GPU 显存")
log("=" * 60)
script_path = Path(__file__).resolve()
t_start = time.time()
successful = 0
failed = 0
for idx in range(args.num_samples):
log(f" [{idx + 1}/{args.num_samples}] 启动 worker {idx} ...")
t0 = time.time()
try:
proc = subprocess.run(
[
sys.executable, str(script_path),
"--worker", str(idx),
"--output_dir", str(output_dir),
"--dataset", args.dataset,
"--temperature", str(args.temperature),
"--seed", str(args.seed),
"--diffusion_snr", str(args.diffusion_snr),
],
capture_output=True,
text=True,
timeout=600, # 单个晶体最多 10 分钟
)
except subprocess.TimeoutExpired:
failed += 1
log(f" 超时 (>600s)")
continue
elapsed = time.time() - t0
stdout = proc.stdout.strip()
stderr = proc.stderr.strip()
# 解析输出
if proc.returncode == 0 and f"OK:{idx}" in stdout:
successful += 1
# 提取原子数信息
atom_info = ""
for line in stdout.split("\n"):
if f"OK:{idx}" in line:
atom_info = line.split(":", 2)[-1] if ":" in line else ""
break
log(
f" 成功 ({elapsed:.0f}s) {atom_info}"
f" [OK={successful} FAIL={failed}]"
)
else:
failed += 1
# 提取错误信息
fail_reason = "unknown"
for line in stdout.split("\n"):
if f"FAIL:{idx}" in line:
fail_reason = line.split(f"FAIL:{idx}:")[-1]
break
if not fail_reason or fail_reason == "unknown":
err_tail = stderr[-300:] if stderr else ""
fail_reason = err_tail.replace("\n", " | ")
log(f" 失败 ({elapsed:.0f}s): {fail_reason}")
# 汇总
t_elapsed = time.time() - t_start
log("")
log("=" * 60)
log(f"生成完成!")
log(f" 成功: {successful}/{args.num_samples}")
log(f" 失败: {failed}/{args.num_samples}")
log(f" 总耗时: {t_elapsed:.0f}s ({t_elapsed / 60:.1f}min)")
log(f" 平均每晶体: {t_elapsed / max(args.num_samples, 1):.0f}s")
log("=" * 60)
# 保存元数据
cif_files = list(cifs_dir.glob("gen_*.cif"))
meta = {
"framework": "pd",
"dataset": args.dataset,
"num_requested": args.num_samples,
"num_successful": successful,
"num_failed": failed,
"num_saved_cifs": len(cif_files),
"temperature": args.temperature,
"seed": args.seed,
"diffusion_snr": args.diffusion_snr,
"time_seconds": round(t_elapsed, 1),
"avg_per_crystal_seconds": round(
t_elapsed / max(args.num_samples, 1), 1
),
"timestamp": time.strftime("%Y-%m-%d %H:%M:%S"),
}
meta_path = output_dir / "metadata.json"
with open(meta_path, "w") as f:
json.dump(meta, f, indent=2, ensure_ascii=False)
log(f"元数据: {meta_path}")
# ---------------------------------------------------------------------------
# 入口
# ---------------------------------------------------------------------------
def main():
args = parse_args()
if args.worker is not None:
# Worker 模式: 在独立子进程中生成单个晶体
worker_main(args)
else:
# 调度模式: 逐个启动子进程
scheduler_main(args)
if __name__ == "__main__":
main()3.3.3 采样指标合并对比移植原版的评估指标Python代码,点击后展开脚本#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
SGEquiDiff 生成评估指标模块。
从 sgequidiff原版/src/utils/eval_utils.py 移植关键评估函数,
"""
from __future__ import annotations
import copy
import itertools
import json
import os
import warnings
from collections import Counter
from pathlib import Path
from typing import Dict, List, Optional, Sequence, Tuple
from zipfile import ZipFile
from tempfile import TemporaryDirectory
import numpy as np
from pymatgen.analysis.structure_matcher import StructureMatcher
from pymatgen.core import Composition, Element, Structure, Lattice
from pymatgen.io.ase import AseAtomsAdaptor
from scipy.spatial.distance import jensenshannon
from scipy.stats import wasserstein_distance
from smact.screening import smact_filter
from tqdm import tqdm
import ase.io
# ============================================================
# 常量
# ============================================================
# MP-20 训练集空间群分布(用于JSD对比)
MP20_SPACE_GROUP_DIST = None
# MP-20 训练集 Wyckoff 维度分布
MP20_WYCKOFF_DIM_DIST = None
# MP-20 训练集密度统计
MP20_DENSITY_STATS = None
# ============================================================
# 晶体有效性检查
# ============================================================
def smact_validity(composition: Composition) -> bool:
"""SMACT 化学有效性检查。
检查元素的化合价是否能平衡。
"""
try:
comp_dict = composition.to_reduced_dict
elements = [Element(el) for el in comp_dict.keys()]
amounts = list(comp_dict.values())
# smact_filter 返回 (valid, reason)
return smact_filter(elements, amounts)[0]
except Exception:
return False
def structure_validity(structure: Structure) -> bool:
"""结构有效性检查。
标准:
- 任意两个原子间距 > 0.5 Angstrom
- 晶胞体积 > 0.1 Angstrom^3
- 原子数量 > 0
"""
if len(structure) == 0:
return False
if structure.volume < 0.1:
return False
# 检查最小原子间距
dist_matrix = structure.distance_matrix
np.fill_diagonal(dist_matrix, np.inf)
min_dist = dist_matrix.min()
if min_dist < 0.5:
return False
return True
def compute_validity(structures: Sequence[Structure]) -> Dict:
"""计算有效性指标。
Args:
structures: pymatgen Structure 列表
Returns:
{
"total": int,
"valid_structure": int,
"valid_composition": int,
"valid_both": int,
"structure_validity_ratio": float,
"composition_validity_ratio": float,
"validity_ratio": float,
}
"""
total = len(structures)
n_struct_valid = 0
n_comp_valid = 0
n_valid = 0
for s in structures:
struct_ok = structure_validity(s)
comp_ok = smact_validity(s.composition)
if struct_ok:
n_struct_valid += 1
if comp_ok:
n_comp_valid += 1
if struct_ok and comp_ok:
n_valid += 1
return {
"total": total,
"valid_structure": n_struct_valid,
"valid_composition": n_comp_valid,
"valid_both": n_valid,
"structure_validity_ratio": n_struct_valid / max(total, 1),
"composition_validity_ratio": n_comp_valid / max(total, 1),
"validity_ratio": n_valid / max(total, 1),
}
# ============================================================
# 结构指纹与多样性
# ============================================================
def get_structure_fingerprint(structure: Structure) -> np.ndarray:
"""计算结构指纹(加权原子位置直方图)。"""
try:
from matminer.featurizers.structure import SiteStatsFingerprint
from matminer.featurizers.site import CrystalNNFingerprint
nn_fp = CrystalNNFingerprint.from_preset("ops")
site_fp = SiteStatsFingerprint(nn_fp)
return site_fp.featurize(structure)
except Exception:
return np.zeros(256)
def get_fingerprint_pairwise_dist(fp_array: np.ndarray) -> float:
"""计算指纹矩阵中所有配对之间的平均距离。"""
n = len(fp_array)
if n <= 1:
return 0.0
dists = []
for i in range(n):
for j in range(i + 1, n):
d = np.linalg.norm(fp_array[i] - fp_array[j])
dists.append(d)
return float(np.mean(dists)) if dists else 0.0
def compute_diversity(structures: Sequence[Structure]) -> Dict:
"""计算多样性指标。
Returns:
{
"structural_diversity": float,
"composition_diversity": float,
}
"""
if len(structures) <= 1:
return {"structural_diversity": 0.0, "composition_diversity": 0.0}
# 结构指纹多样性 - 用固定长度填充处理不一致的指纹维度
fingerprints = []
for s in tqdm(structures, desc="结构指纹", leave=False):
fp = get_structure_fingerprint(s)
fp = np.asarray(fp).flatten()
fingerprints.append(fp)
# 填充到相同长度
max_len = max(len(fp) for fp in fingerprints)
fp_padded = np.zeros((len(fingerprints), max_len))
for i, fp in enumerate(fingerprints):
fp_padded[i, :len(fp)] = fp
struct_div = get_fingerprint_pairwise_dist(fp_padded)
# 组成多样性(基于元素数量向量)
comp_vectors = []
for s in tqdm(structures, desc="元素向量", leave=False):
comp = s.composition
vec = np.zeros(100)
for el, amt in comp.to_reduced_dict.items():
try:
element = Element(el)
z = element.Z
if 1 <= z < 100:
vec[z] = amt
except Exception:
continue
comp_vectors.append(vec)
comp_array = np.array(comp_vectors)
comp_div = get_fingerprint_pairwise_dist(comp_array)
return {"structural_diversity": struct_div, "composition_diversity": comp_div}
# ============================================================
# 覆盖率(与参考集对比)
# ============================================================
def compute_coverage(
gen_structures: Sequence[Structure],
ref_structures: Sequence[Structure],
struc_cutoff: float = 0.4,
comp_cutoff: float = 10.0,
) -> Dict:
"""计算覆盖率指标(生成结构中有多少比例能在参考集中找到匹配)。
Args:
gen_structures: 生成的结构列表
ref_structures: 参考结构列表(如训练集)
struc_cutoff: 结构匹配 cutoff
comp_cutoff: 组成匹配 cutoff
Returns:
{"coverage": float, "num_matched": int, "total": int}
"""
matcher = StructureMatcher(
stol=struc_cutoff,
angle_tol=comp_cutoff,
ltol=comp_cutoff,
)
num_matched = 0
total = len(gen_structures)
# 对每个生成结构,检查是否能在参考集中找到匹配
for gen_s in tqdm(gen_structures, desc="覆盖率", leave=False):
if total == 0:
break
matched = False
for ref_s in ref_structures:
if matcher.fit(gen_s, ref_s):
matched = True
break
if matched:
num_matched += 1
return {
"coverage": num_matched / max(total, 1),
"num_matched": num_matched,
"total": total,
}
# ============================================================
# 分布指标
# ============================================================
def get_space_group_number(structure: Structure) -> int:
"""获取空间群编号(1-230)。"""
try:
from pymatgen.symmetry.analyzer import SpacegroupAnalyzer
sga = SpacegroupAnalyzer(structure)
return sga.get_space_group_number()
except Exception:
return 1
def compute_space_group_jsd(
gen_structures: Sequence[Structure],
ref_distribution: Optional[np.ndarray] = None,
) -> float:
"""计算空间群分布的 Jensen-Shannon Divergence。
Args:
gen_structures: 生成结构列表
ref_distribution: 参考分布(230维),默认均匀分布
Returns:
JSD 值(越小越接近)
"""
# 计算生成结构的空间群分布
sg_counts = Counter()
for s in tqdm(gen_structures, desc="空间群", leave=False):
sg = get_space_group_number(s)
sg_counts[sg] += 1
gen_dist = np.zeros(230)
for sg, count in sg_counts.items():
if 1 <= sg <= 230:
gen_dist[sg - 1] = count
gen_dist = gen_dist / max(gen_dist.sum(), 1)
# 参考分布
if ref_distribution is not None:
ref_dist = np.array(ref_distribution)
else:
ref_dist = np.ones(230) / 230 # 均匀分布
# 计算 JSD
return float(jensenshannon(gen_dist, ref_dist))
def compute_density_wasserstein(
gen_structures: Sequence[Structure],
ref_densities: Optional[List[float]] = None,
) -> float:
"""计算密度分布的 Wasserstein 距离。"""
gen_densities = []
for s in tqdm(gen_structures, desc="密度", leave=False):
try:
gen_densities.append(s.density)
except Exception:
pass
if not gen_densities:
return 1.0
if ref_densities is not None and len(ref_densities) > 0:
return float(wasserstein_distance(gen_densities, ref_densities))
else:
return 0.0
def compute_num_elements_wasserstein(
gen_structures: Sequence[Structure],
ref_num_elements: Optional[List[int]] = None,
) -> float:
"""计算元素数量分布的 Wasserstein 距离。"""
gen_counts = []
for s in gen_structures:
gen_counts.append(len(s.composition.to_reduced_dict))
if ref_num_elements is not None and len(ref_num_elements) > 0:
return float(wasserstein_distance(gen_counts, ref_num_elements))
else:
return 0.0
def compute_distribution_metrics(
gen_structures: Sequence[Structure],
ref_structures: Optional[Sequence[Structure]] = None,
) -> Dict:
"""计算分布指标。
Returns:
{
"space_group_jsd": float,
"density_wasserstein": float,
"num_elements_wasserstein": float,
}
"""
# 计算参考分布的统计量
ref_sg_dist = None
ref_densities = None
ref_num_elements = None
if ref_structures is not None and len(ref_structures) > 0:
sg_counts = Counter()
densities = []
num_elems = []
for s in tqdm(ref_structures, desc="参考分布", leave=False):
sg = get_space_group_number(s)
sg_counts[sg] += 1
try:
densities.append(s.density)
except Exception:
pass
num_elems.append(len(s.composition.to_reduced_dict))
ref_sg_dist = np.zeros(230)
for sg, count in sg_counts.items():
if 1 <= sg <= 230:
ref_sg_dist[sg - 1] = count
ref_sg_dist = ref_sg_dist / max(ref_sg_dist.sum(), 1)
ref_densities = densities
ref_num_elements = num_elems
return {
"space_group_jsd": compute_space_group_jsd(gen_structures, ref_sg_dist),
"density_wasserstein": compute_density_wasserstein(gen_structures, ref_densities),
"num_elements_wasserstein": compute_num_elements_wasserstein(gen_structures, ref_num_elements),
}
# ============================================================
# 结构加载
# ============================================================
def load_structures_from_npz(path: str) -> List[Structure]:
"""从 .npz 文件加载晶体结构(SGEquiDiff 格式)。"""
pass # 需要根据实际 npz 格式实现
def load_structures_from_cif_dir(path: str) -> List[Structure]:
"""从 CIF 目录加载结构。"""
structures = []
path = Path(path)
if not path.exists():
return structures
for fname in sorted(path.iterdir()):
if fname.suffix in (".cif",):
try:
from pymatgen.io.cif import CifParser
parser = CifParser(str(fname), occupancy_tolerance=100)
s = parser.parse_structures()[0]
structures.append(s)
except Exception as e:
warnings.warn(f"读取 {fname.name} 失败: {e}")
return structures
def load_structures_from_json(path: str) -> List[Structure]:
"""从 JSON 文件加载结构。"""
structures = []
path = Path(path)
if not path.exists():
return structures
with open(path) as f:
data = json.load(f)
for item in data:
try:
lattice = Lattice.from_parameters(
*item["lattice_params"]
)
species = item["species"]
coords = item["frac_coords"]
s = Structure(lattice, species, coords, coords_are_cartesian=False)
structures.append(s)
except Exception as e:
warnings.warn(f"解析结构失败: {e}")
return structures
# ============================================================
# 完整评估
# ============================================================
def evaluate_structures(
structures: Sequence[Structure],
ref_structures: Optional[Sequence[Structure]] = None,
dataset_name: str = "mp_20",
) -> Dict:
"""对结构列表进行全面的评估。
Args:
structures: 要评估的结构列表
ref_structures: 参考结构列表(训练集),用于覆盖率和分布指标
dataset_name: 数据集名称
Returns:
包含所有评估指标的字典
"""
if len(structures) == 0:
return {"error": "没有结构可供评估"}
metrics = {}
metrics["num_structures"] = len(structures)
# 1. 有效性
print("计算有效性指标...")
metrics["validity"] = compute_validity(structures)
# 2. 多样性
print("计算多样性指标...")
metrics["diversity"] = compute_diversity(structures)
# 3. 覆盖率(需要参考集)
if ref_structures is not None and len(ref_structures) > 0:
print("计算覆盖率...")
metrics["coverage"] = compute_coverage(structures, ref_structures)
else:
metrics["coverage"] = {"coverage": 0.0, "num_matched": 0, "total": len(structures)}
# 4. 分布指标
print("计算分布指标...")
metrics["distribution"] = compute_distribution_metrics(structures, ref_structures)
return metrics
def compute_comparison(metrics_pt: Dict, metrics_pd: Dict) -> Dict:
"""对比 PT 和 PD 的评估指标,计算相对误差。
Args:
metrics_pt: PyTorch 评估结果
metrics_pd: Paddle 评估结果
Returns:
{
"comparison": {指标名: {"pt": x, "pd": y, "abs_diff": d, "rel_diff": d/|x|}},
"all_within_5pct": bool
}
"""
comparison = {}
def _compare(key: str, pt_val: float, pd_val: float):
abs_diff = abs(pt_val - pd_val)
rel_diff = abs_diff / max(abs(pt_val), 1e-10)
return {
"pt": float(pt_val),
"pd": float(pd_val),
"abs_diff": float(abs_diff),
"rel_diff": float(rel_diff),
"within_5pct": rel_diff < 0.05,
}
# 有效性指标
for k in ("validity_ratio", "structure_validity_ratio", "composition_validity_ratio"):
pt_v = metrics_pt.get("validity", {}).get(k, 0)
pd_v = metrics_pd.get("validity", {}).get(k, 0)
comparison[f"validity.{k}"] = _compare(k, pt_v, pd_v)
# 多样性指标
for k in ("structural_diversity", "composition_diversity"):
pt_v = metrics_pt.get("diversity", {}).get(k, 0)
pd_v = metrics_pd.get("diversity", {}).get(k, 0)
comparison[f"diversity.{k}"] = _compare(k, pt_v, pd_v)
# 覆盖率
for k in ("coverage",):
pt_v = metrics_pt.get("coverage", {}).get(k, 0)
pd_v = metrics_pd.get("coverage", {}).get(k, 0)
comparison[f"coverage.{k}"] = _compare(k, pt_v, pd_v)
# 分布指标
for k in ("space_group_jsd", "density_wasserstein", "num_elements_wasserstein"):
pt_v = metrics_pt.get("distribution", {}).get(k, 0)
pd_v = metrics_pd.get("distribution", {}).get(k, 0)
comparison[f"distribution.{k}"] = _compare(k, pt_v, pd_v)
# 检查是否所有指标都在 5% 以内
all_within = all(v.get("within_5pct", False) for v in comparison.values())
return {
"comparison": comparison,
"all_within_5pct": all_within,
"num_metrics": len(comparison),
"num_passed": sum(1 for v in comparison.values() if v.get("within_5pct", False)),
}
def print_comparison_summary(comparison_result: Dict):
"""打印对比结果摘要。"""
print("\n" + "=" * 70)
print("PT vs PD 采样指标对比报告")
print("=" * 70)
data = comparison_result["comparison"]
print(f"\n{'指标':<40s} {'PT':>10s} {'PD':>10s} {'|diff|':>12s} {'rel_diff':>10s} {'结果':>8s}")
print("-" * 90)
for key, vals in sorted(data.items()):
status = "PASS" if vals["within_5pct"] else "FAIL"
print(
f"{key:<40s} {vals['pt']:>10.4f} {vals['pd']:>10.4f} "
f"{vals['abs_diff']:>12.6f} {vals['rel_diff']:>10.4f} {status:>8s}"
)
print("-" * 90)
print(f"\n总指标数: {comparison_result['num_metrics']}")
print(f"通过数: {comparison_result['num_passed']}")
print(f"整体结果: {'PASS (全部在 5% 以内)' if comparison_result['all_within_5pct'] else 'FAIL (部分指标超出 5%)'}")
print("=" * 70)开始执行对比Python脚本,点击后展开脚本#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
晶体采样指标评估与 PT/PD 对比脚本。
使用流程:
步骤1 (PT): cd sgequidiff原版代码目录 && source .venv/bin/activate
python scripts/generate_samples_pt.py ...
步骤2 (PD): conda activate ppmat
python scripts/generate_samples_pd.py ...
步骤3 (对比): 在本脚本中指定 PT/PD 输出目录即可对比
"""
import argparse
import json
import os
import sys
import time
from pathlib import Path
from typing import Dict, List, Optional
import numpy as np
def parse_args():
parser = argparse.ArgumentParser(description="采样指标评估与对比")
# 输入目录
parser.add_argument("--pt_cif_dir", type=str, default=None,
help="PT 生成晶体的 CIF 目录")
parser.add_argument("--pd_cif_dir", type=str, default=None,
help="PD 生成晶体的 CIF 目录")
parser.add_argument("--ref_cif_dir", type=str, default=None,
help="参考集 CIF 目录(如训练集)")
# 输出
parser.add_argument("--output_dir", type=str, default=None,
help="评估结果输出目录")
parser.add_argument("--dataset", type=str, default="mp_20",
choices=["mp_20", "mpts_52"])
parser.add_argument("--num_ref", type=int, default=None,
help="参考集使用数量,默认全部")
return parser.parse_args()
def main():
args = parse_args()
# 确定输出目录
if args.output_dir is None:
timestamp = time.strftime("%Y%m%d_%H%M%S")
output_dir = Path(f"./outputs/sampling_eval_{timestamp}")
else:
output_dir = Path(args.output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
# 日志
log_path = output_dir / "evaluation.log"
log_file = open(log_path, "w")
def log(msg):
print(msg, flush=True)
log_file.write(msg + "\n")
log_file.flush()
log(f"{'='*70}")
log(f"SGEquiDiff 采样指标评估")
log(f"{'='*70}")
log(f"时间: {time.strftime('%Y-%m-%d %H:%M:%S')}")
log(f"数据集: {args.dataset}")
log(f"输出目录: {output_dir}")
log("")
# 导入评估模块
sys.path.insert(0, str(Path(__file__).resolve().parent))
from generation_metrics import (
load_structures_from_cif_dir,
load_structures_from_json,
evaluate_structures,
compute_comparison,
print_comparison_summary,
)
# 加载 PT 结构
pt_structures = []
if args.pt_cif_dir:
pt_path = Path(args.pt_cif_dir)
if pt_path.exists():
log(f"加载 PT 结构: {pt_path}")
pt_structures = load_structures_from_cif_dir(str(pt_path))
log(f" PT 结构数: {len(pt_structures)}")
else:
log(f" PT 目录不存在: {pt_path}")
# 加载 PD 结构
pd_structures = []
if args.pd_cif_dir:
pd_path = Path(args.pd_cif_dir)
if pd_path.exists():
log(f"加载 PD 结构: {pd_path}")
pd_structures = load_structures_from_cif_dir(str(pd_path))
log(f" PD 结构数: {len(pd_structures)}")
else:
log(f" PD 目录不存在: {pd_path}")
# 加载参考结构
ref_structures = []
if args.ref_cif_dir:
ref_path = Path(args.ref_cif_dir)
if ref_path.exists():
log(f"加载参考结构: {ref_path}")
all_ref = load_structures_from_cif_dir(str(ref_path))
if args.num_ref:
ref_structures = all_ref[:args.num_ref]
else:
ref_structures = all_ref
log(f" 参考结构数: {len(ref_structures)}")
else:
log(f" 参考目录不存在: {ref_path}")
# 评估 PT
metrics_pt = None
if pt_structures:
log(f"\n{'='*70}")
log("评估 PT 生成样本...")
log(f"{'='*70}")
metrics_pt = evaluate_structures(pt_structures, ref_structures, args.dataset)
_save_metrics(output_dir / "metrics_pt.json", metrics_pt)
log(f"\nPT 评估结果已保存到: {output_dir / 'metrics_pt.json'}")
else:
log("\n跳过 PT 评估 (无结构)")
# 评估 PD
metrics_pd = None
if pd_structures:
log(f"\n{'='*70}")
log("评估 PD 生成样本...")
log(f"{'='*70}")
metrics_pd = evaluate_structures(pd_structures, ref_structures, args.dataset)
_save_metrics(output_dir / "metrics_pd.json", metrics_pd)
log(f"\nPD 评估结果已保存到: {output_dir / 'metrics_pd.json'}")
else:
log("\n跳过 PD 评估 (无结构)")
# 对比
comparison = None
if metrics_pt and metrics_pd:
log(f"\n{'='*70}")
log("PT vs PD 对比")
log(f"{'='*70}")
comparison = compute_comparison(metrics_pt, metrics_pd)
print_comparison_summary(comparison)
# 保存对比结果
with open(output_dir / "comparison.json", "w") as f:
# 转换 numpy 类型
json.dump(comparison, f, indent=2, default=_json_default)
log(f"\n对比结果已保存到: {output_dir / 'comparison.json'}")
elif pt_structures:
log("\n仅 PT 评估完成 (无 PD 数据无法对比)")
elif pd_structures:
log("\n仅 PD 评估完成 (无 PT 数据无法对比)")
else:
log("\n未加载任何结构,评估跳过")
# 汇总
log(f"\n{'='*70}")
log("评估汇总")
log(f"{'='*70}")
if comparison:
log(f" 总指标数: {comparison['num_metrics']}")
log(f" 通过数: {comparison['num_passed']}")
log(f" 结果: {'全部通过' if comparison['all_within_5pct'] else '部分未通过'}")
log(f" 输出目录: {output_dir}")
log(f" 日志文件: {log_path}")
log(f"{'='*70}")
log_file.close()
def _save_metrics(path, metrics):
"""保存评估指标到 JSON。"""
with open(path, "w") as f:
json.dump(metrics, f, indent=2, default=_json_default)
def _json_default(obj):
"""JSON 序列化的默认处理器。"""
if isinstance(obj, (np.integer,)):
return int(obj)
if isinstance(obj, (np.floating,)):
return float(obj)
if isinstance(obj, (np.ndarray,)):
return obj.tolist()
return str(obj)
if __name__ == "__main__":
main() |
learncat163
marked this pull request as draft
May 8, 2026 07:39
learncat163
force-pushed
the
feature-add-sgequidiff
branch
from
May 8, 2026 11:09
973376a to
841a80e
Compare
learncat163
marked this pull request as ready for review
May 8, 2026 11:13
Collaborator
|
#258 辛苦参考该文档和套件已有的文件组织格式 |
Collaborator
|
代码里有大量 print 调试输出 |
learncat163
force-pushed
the
feature-add-sgequidiff
branch
from
June 24, 2026 08:02
841a80e to
02f039d
Compare
|
ppmat-bot seems not to be a GitHub user. You need a GitHub account to be able to sign the CLA. If you have already a GitHub account, please add the email address used for this commit to your account. You have signed the CLA already but the status is still pending? Let us recheck it. |
learncat163
marked this pull request as draft
June 24, 2026 08:09
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
【已经废弃】新的PR #300
SGEquiDiff 迁移
1. 概述
将 SGEquiDiff(Symmetry-Guided Equivariant Diffusion)晶体生成模型从 PyTorch 框架迁移至 Paddle。
SGEquiDiff 是一种基于空间群对称性约束的等变扩散模型,用于生成符合晶体学对称性要求的晶体结构。该模型涉及复杂的对称群操作、非对称单元(ASU)坐标变换、Wyckoff 位置采样等计算。
项目的运行、训练、采样,在星河项目:https://aistudio.baidu.com/projectdetail/10304319
2. 权重信息
2.1 原版权重
原始的权重在:https://drive.google.com/drive/folders/1ONwO53i6oG1_yBqP0zPQ-IV_zLoIWyQR 如果不方便访问google drive,可以使用aistudio上的原版搬运镜像:https://aistudio.baidu.com/modelsdetail/47157/space
2.2 转换权重
基于2种数据集 MP-20 数据集 和 MPTS-52 数据集 的权重,分别进行转换,地址:https://aistudio.baidu.com/modelsdetail/47158?modelId=47158
其中,paddle的http在线下载地址内置在
ppmat/models/sgequidiff/constants.py程序默认会自动从http上下载权重到合适位置。3. 精度对齐验证
3.1 前向精度对齐
前向精度对齐主要考虑4部分:
对比 PT/PD 4个子模块的前向传播输出,针对确定性输出(diffusion, space_group, lattice 非随机部分)计算最大绝对误差和平均绝对误差。
3.1.1 扩散模型(Diffusion)对齐
3.1.2 空间群模型(log_prob)
3.1.3 晶格模型(Lattice)
3.1.4 Wyckoff 模型
Wyckoff 深度依赖随机数发生器RNG,导致难易对齐,除非对PT的原版代码直接进行hook修改,这有一定的破坏性,暂时无法实现。
前向精度对齐命令Python脚本,点击后展开脚本
3.2 反向精度对齐
使用固定 crystal 输入 + 固定 target score 进行 AdamW 微训练,对比 PT/PD loss 轨迹和梯度范数趋势。
non_equivariant_drift_model.mlp_out3.2.1 反向对齐loss对比
反向对齐loss的Python脚本,点击后展开脚本
3.3 采样指标对齐
采样指标对齐,在下一个评论里;github禁止PR超过64K,拆分2段!!!