Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions ppmat/datasets/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@
from ppmat.datasets.msd_nmr_dataset import MSDnmrDataset
from ppmat.datasets.msd_nmr_dataset import MSDnmrinfos
from ppmat.datasets.density_dataset import DensityDataset
from ppmat.datasets.dm2_dataset import DM2StructureDataset
from ppmat.datasets.small_density_dataset import SmallDensityDataset
from ppmat.datasets.num_atom_crystal_dataset import NumAtomsCrystalDataset
from ppmat.datasets.oc20_s2ef_dataset import OC20S2EFDataset # noqa
Expand All @@ -65,6 +66,7 @@
"MSDnmrDataset",
"MatbenchDataset",
"DensityDataset",
"DM2StructureDataset",
"SmallDensityDataset",
"OMol25Dataset",
]
Expand Down
181 changes: 181 additions & 0 deletions ppmat/datasets/dm2_dataset.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,181 @@
# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved.

# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at

# http://www.apache.org/licenses/LICENSE-2.0

# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from __future__ import annotations

import glob
from typing import Dict
from typing import List
from typing import Optional
from typing import Sequence
from typing import Union

import ase.io
import numpy as np
import paddle
from ase.neighborlist import primitive_neighbor_list
from paddle.io import Dataset

from ppmat.datasets.geometric_data_type.data import Data


def _expand_paths(paths: Union[str, Sequence[str]]) -> List[str]:
if isinstance(paths, str):
paths = [paths]

expanded = []
for path in paths:
matches = sorted(glob.glob(path))
if matches:
expanded.extend(matches)
else:
expanded.append(path)
return expanded


def _as_pbc_tuple(pbc) -> tuple[bool, bool, bool]:
arr = np.asarray(pbc, dtype=bool).reshape(-1)
if arr.size == 1:
arr = np.repeat(arr, 3)
return tuple(bool(x) for x in arr[:3])


def ase_atoms_to_dm2_data(
atoms,
species_to_index: Dict[int, int],
cutoff: float,
cooling_rate: Optional[float] = None,
) -> Data:
"""Convert an ASE ``Atoms`` object to PaddleMaterials geometric ``Data``."""

atomic_numbers = np.asarray(atoms.numbers, dtype=np.int64)
x = np.asarray([species_to_index[int(z)] for z in atomic_numbers], dtype=np.int64)
pbc = _as_pbc_tuple(atoms.pbc)
cell = np.asarray(atoms.cell.array, dtype=np.float32)
positions = np.asarray(atoms.positions, dtype=np.float32)

src, dst, disp = primitive_neighbor_list(
"ijD",
cutoff=cutoff,
pbc=pbc,
cell=cell,
positions=positions,
numbers=atomic_numbers,
)
data = Data(
x=paddle.to_tensor(x, dtype="int64"),
pos=paddle.to_tensor(positions, dtype="float32"),
edge_index=paddle.to_tensor(np.stack([src, dst]), dtype="int64"),
edge_attr=paddle.to_tensor(disp, dtype="float32"),
atomic_numbers=paddle.to_tensor(atomic_numbers, dtype="int64"),
lattice=paddle.to_tensor(cell[None, :, :], dtype="float32"),
pbc=paddle.to_tensor(np.asarray(pbc, dtype=bool)[None, :], dtype="bool"),
num_nodes=len(atomic_numbers),
)
if cooling_rate is not None:
data.cooling_rate = paddle.to_tensor([cooling_rate], dtype="float32")
return data


class DM2StructureDataset(Dataset):
"""ASE-backed dataset for DM2 disordered-material denoising.

Args:
paths (str|Sequence[str]): Structure files or glob patterns. ASE is used for
reading, so formats such as ``lammps-data``, ``extxyz`` and ``cif`` are
supported through ``file_format``.
file_format (Optional[str]): ASE format string. Leave ``None`` to let ASE
infer the format.
cutoff (float): Large neighbor cutoff used before rattle/downselect.
species (Optional[Sequence[int]]): Ordered atomic numbers for species
encoding. If omitted, the dataset infers a sorted species list.
duplicate (int): Repeat each loaded structure this many times per epoch.
cooling_rates (Optional[float|Sequence[float]]): Per-structure conditioning
values. A scalar or one-item sequence is broadcast to all structures.
DM2 conditional training conventionally uses ``log10(cooling_rate)``;
set ``log10_cooling_rate=True`` to apply that transform here.
log10_cooling_rate (bool): Whether to store log10-transformed cooling rates.
"""

def __init__(
self,
paths: Union[str, Sequence[str]],
file_format: Optional[str] = None,
cutoff: float = 10.0,
species: Optional[Sequence[int]] = None,
duplicate: int = 128,
cooling_rates: Optional[Sequence[float]] = None,
log10_cooling_rate: bool = True,
):
super().__init__()
self.paths = _expand_paths(paths)
if len(self.paths) == 0:
raise ValueError("DM2StructureDataset received no structure files.")
self.file_format = file_format
self.cutoff = cutoff
self.duplicate = int(duplicate)
if self.duplicate <= 0:
raise ValueError("duplicate must be a positive integer.")

atoms_list = [
ase.io.read(path, format=file_format)
for path in self.paths
]

if species is None:
unique_species = sorted(
{
int(number)
for atoms in atoms_list
for number in np.asarray(atoms.numbers).tolist()
}
)
else:
unique_species = [int(number) for number in species]
self.species = unique_species
self.species_to_index = {z: idx for idx, z in enumerate(unique_species)}

if cooling_rates is not None:
if isinstance(cooling_rates, (int, float)):
cooling_rates = [float(cooling_rates)] * len(atoms_list)
elif len(cooling_rates) == 1 and len(atoms_list) > 1:
cooling_rates = [float(cooling_rates[0])] * len(atoms_list)
elif len(cooling_rates) != len(atoms_list):
raise ValueError(
"cooling_rates must be omitted, a scalar, a single value to "
"broadcast, or have the same length as paths."
)

self.graphs = []
for idx, atoms in enumerate(atoms_list):
cooling_rate = None
if cooling_rates is not None:
cooling_rate = float(cooling_rates[idx])
if log10_cooling_rate:
cooling_rate = float(np.log10(cooling_rate))
self.graphs.append(
ase_atoms_to_dm2_data(
atoms=atoms,
species_to_index=self.species_to_index,
cutoff=cutoff,
cooling_rate=cooling_rate,
)
)

def __len__(self):
return len(self.graphs) * self.duplicate

def __getitem__(self, idx):
graph = self.graphs[idx % len(self.graphs)]
return graph.clone()
2 changes: 2 additions & 0 deletions ppmat/metrics/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,11 +17,13 @@
import paddle # noqa

from ppmat.metrics.csp_metric import CSPMetric
from ppmat.metrics.dm2_metric import DM2AmorphousGenerationMetric
from ppmat.metrics.diffnmr_streaming_adapter import DiffNMRStreamingAdapter

__all__ = [
"build_metric",
"CSPMetric",
"DM2AmorphousGenerationMetric",
"DiffNMRStreamingAdapter",
# "DiffNMRMetric",
# "NLL", "CrossEntropyMetric", "SumExceptBatchMetric", "SumExceptBatchKL",
Expand Down
156 changes: 156 additions & 0 deletions ppmat/metrics/dm2_metric.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,156 @@
# Copyright (c) 2025 PaddlePaddle Authors. All Rights Reserved.

# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at

# http://www.apache.org/licenses/LICENSE-2.0

# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

from __future__ import annotations

import glob
from typing import Dict
from typing import List
from typing import Sequence
from typing import Union

import ase.io
import numpy as np
from ase import Atoms


def _expand_paths(paths: Union[str, Sequence[str]]) -> List[str]:
if isinstance(paths, str):
paths = [paths]

expanded = []
for pattern in paths:
matches = sorted(glob.glob(pattern))
if not matches:
raise FileNotFoundError(f"No DM2 metric reference files matched: {pattern}")
expanded.extend(matches)
return expanded


def _structure_array_to_atoms(structure_array: Dict) -> Atoms:
frac_coords = np.asarray(structure_array["frac_coords"], dtype=np.float64)
atom_types = np.asarray(structure_array["atom_types"], dtype=np.int64)
lattice = np.asarray(structure_array["lattice"], dtype=np.float64).reshape(3, 3)
return Atoms(
numbers=atom_types,
scaled_positions=frac_coords,
cell=lattice,
pbc=True,
)


def _rdf_histogram(atoms_list: Sequence[Atoms], cutoff: float, bins: int):
hist = np.zeros(bins, dtype=np.float64)
for atoms in atoms_list:
distances = atoms.get_all_distances(mic=True)
upper = distances[np.triu_indices_from(distances, k=1)]
upper = upper[(upper > 1e-8) & (upper <= cutoff)]
values, _ = np.histogram(upper, bins=bins, range=(0.0, cutoff))
hist += values.astype(np.float64)
total = hist.sum()
if total > 0:
hist /= total
return hist


def _wasserstein_1d_from_hist(hist_a, hist_b, cutoff: float):
bin_width = cutoff / len(hist_a)
return float(np.abs(np.cumsum(hist_a) - np.cumsum(hist_b)).sum() * bin_width)


def _mean_coordination(
atoms_list: Sequence[Atoms],
center_atomic_number: int,
neighbor_atomic_number: int,
cutoff: float,
):
values = []
for atoms in atoms_list:
numbers = np.asarray(atoms.numbers, dtype=np.int64)
center_mask = numbers == int(center_atomic_number)
neighbor_mask = numbers == int(neighbor_atomic_number)
if not np.any(center_mask):
continue
distances = atoms.get_all_distances(mic=True)
valid = (
(distances <= float(cutoff))
& (distances > 1e-8)
& neighbor_mask[None, :]
)
values.extend(valid[center_mask].sum(axis=1).tolist())
if not values:
return float("nan")
return float(np.mean(values))


class DM2AmorphousGenerationMetric:
"""RDF and coordination metrics for DM2 amorphous structure sampling."""

def __init__(
self,
reference_paths: Union[str, Sequence[str]],
file_format: str = "lammps-data",
rdf_cutoff: float = 8.0,
rdf_bins: int = 200,
coordination_cutoffs: Sequence[Dict] = (),
):
self.reference_paths = _expand_paths(reference_paths)
self.file_format = file_format
self.rdf_cutoff = float(rdf_cutoff)
self.rdf_bins = int(rdf_bins)
self.coordination_cutoffs = list(coordination_cutoffs)
self.reference_atoms = [
ase.io.read(path, format=file_format) for path in self.reference_paths
]
self.reference_rdf = _rdf_histogram(
self.reference_atoms,
cutoff=self.rdf_cutoff,
bins=self.rdf_bins,
)

def __call__(self, pred_structures: Sequence[Dict]):
pred_atoms = [_structure_array_to_atoms(item) for item in pred_structures]
pred_rdf = _rdf_histogram(
pred_atoms,
cutoff=self.rdf_cutoff,
bins=self.rdf_bins,
)

metric = {
"rdf_wasserstein": _wasserstein_1d_from_hist(
pred_rdf,
self.reference_rdf,
cutoff=self.rdf_cutoff,
)
}

for item in self.coordination_cutoffs:
name = item.get("name", "coordination")
pred_value = _mean_coordination(
pred_atoms,
center_atomic_number=item["center_atomic_number"],
neighbor_atomic_number=item["neighbor_atomic_number"],
cutoff=item["cutoff"],
)
ref_value = _mean_coordination(
self.reference_atoms,
center_atomic_number=item["center_atomic_number"],
neighbor_atomic_number=item["neighbor_atomic_number"],
cutoff=item["cutoff"],
)
metric[name] = pred_value
metric[f"{name}_reference"] = ref_value
metric[f"{name}_abs_diff"] = abs(pred_value - ref_value)

return metric
4 changes: 4 additions & 0 deletions ppmat/models/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,8 @@
from ppmat.models.common.graph_converter import CrystalNN
from ppmat.models.common.graph_converter import FindPointsInSpheres
from ppmat.models.common.graph_converter import MolecularGraphConverter
from ppmat.models.dm2.dm2 import DM2
from ppmat.models.dm2.dm2 import DM2NequIPDenoiser
from ppmat.models.diffcsp.diffcsp import DiffCSP
from ppmat.models.diffnmr.diffnmr import DiffNMR
from ppmat.models.diffnmr.diffnmr import DiffPrior
Expand All @@ -49,6 +51,8 @@
__all__ = [
"iComformer",
"ComformerGraphConverter",
"DM2",
"DM2NequIPDenoiser",
"DiffCSP",
"FindPointsInSpheres",
"MEGNetPlus",
Expand Down
2 changes: 1 addition & 1 deletion ppmat/models/common/e3nn/nn/_gate.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ def __init__(self, *irreps_outs):
instructions += [tuple(range(i, i + len(irreps_out)))]
i += len(irreps_out)
assert len(irreps_in) == i, (len(irreps_in), i)
irreps_in, p, _ = paddle.sort(x=irreps_in), paddle.argsort(x=irreps_in)
irreps_in, p, _ = irreps_in.sort()
instructions = [tuple(p[i] for i in x) for x in instructions]
self.cut = Extract(irreps_in, self.irreps_outs, instructions)
self.irreps_in = irreps_in.simplify()
Expand Down
Loading