Team Ai
Datasetpublic

OneScience-Group/mp20

MP-20 Dataset Description MP-20 is derived from the Materials Project database (Jain et al., 2013) and contains approximately 45,000 common inorganic materials. It covers most experimentally known materials with no more than 20 atoms in the unit cell. The data includes elemental compositions, CIF structures, space groups, formation energies per atom, DFT band gaps, bulk moduli, and magnetic densities. Source paper: A. Jain et al., Commentary: The Materials Project:… See the full description on the dataset page: https://huggingface.co/datasets/OneScience-Group/mp20.

sourceHugging Facecc-by-4.0updated 2mo agoView on Hugging Face
0likes109downloads
validate_mp20.py155 linesDownload Raw Back to scripts
1#!/usr/bin/env python32"""Validate the standardized OneScience MP-20 dataset package."""3 4from __future__ import annotations5 6import argparse7import csv8import hashlib9import json10import sys11from pathlib import Path12 13import numpy as np14 15 16EXPECTED_SPLITS = {"train": 27136, "val": 9047, "test": 9046}17REQUIRED_COLUMNS = {18    "material_id",19    "formation_energy_per_atom",20    "dft_band_gap",21    "pretty_formula",22    "elements",23    "cif",24    "spacegroup_number",25    "dft_bulk_modulus",26    "dft_mag_density",27}28REQUIRED_ARRAYS = {29    "atomic_numbers.npy",30    "cell.npy",31    "num_atoms.npy",32    "pos.npy",33    "structure_id.npy",34}35PROPERTY_FILES = {36    "formation_energy_per_atom.json",37    "dft_band_gap.json",38    "dft_bulk_modulus.json",39    "dft_mag_density.json",40}41 42 43def require_file(path: Path) -> None:44    if not path.is_file():45        raise FileNotFoundError(f"missing file: {path}")46 47 48def sha256_file(path: Path) -> str:49    digest = hashlib.sha256()50    with path.open("rb") as handle:51        for chunk in iter(lambda: handle.read(1024 * 1024), b""):52            digest.update(chunk)53    return digest.hexdigest()54 55 56def validate_csv(path: Path, expected_rows: int) -> None:57    require_file(path)58    with path.open(encoding="utf-8", newline="") as handle:59        reader = csv.DictReader(handle)60        columns = set(reader.fieldnames or [])61        missing = REQUIRED_COLUMNS - columns62        if missing:63            raise ValueError(f"missing CSV columns in {path}: {sorted(missing)}")64        rows = sum(1 for _ in reader)65    if rows != expected_rows:66        raise ValueError(f"unexpected CSV rows in {path}: {rows} != {expected_rows}")67 68 69def validate_cache(path: Path, expected_structures: int) -> None:70    for name in REQUIRED_ARRAYS | PROPERTY_FILES:71        require_file(path / name)72 73    atomic_numbers = np.load(path / "atomic_numbers.npy", mmap_mode="r", allow_pickle=False)74    cells = np.load(path / "cell.npy", mmap_mode="r", allow_pickle=False)75    num_atoms = np.load(path / "num_atoms.npy", mmap_mode="r", allow_pickle=False)76    positions = np.load(path / "pos.npy", mmap_mode="r", allow_pickle=False)77    structure_ids = np.load(path / "structure_id.npy", mmap_mode="r", allow_pickle=False)78 79    if cells.shape != (expected_structures, 3, 3):80        raise ValueError(f"unexpected cell shape in {path}: {cells.shape}")81    if num_atoms.shape != (expected_structures,):82        raise ValueError(f"unexpected num_atoms shape in {path}: {num_atoms.shape}")83    if structure_ids.shape != (expected_structures,):84        raise ValueError(f"unexpected structure_id shape in {path}: {structure_ids.shape}")85    total_atoms = int(num_atoms.sum())86    if atomic_numbers.shape != (total_atoms,):87        raise ValueError(f"atomic_numbers and num_atoms disagree in {path}")88    if positions.shape != (total_atoms, 3):89        raise ValueError(f"positions and num_atoms disagree in {path}")90 91    for name in PROPERTY_FILES:92        payload = json.loads((path / name).read_text(encoding="utf-8"))93        if len(payload.get("values", [])) != expected_structures:94            raise ValueError(f"unexpected property length in {path / name}")95 96 97def read_checksum_manifest(path: Path) -> list[tuple[str, str]]:98    require_file(path)99    entries: list[tuple[str, str]] = []100    with path.open(encoding="utf-8") as handle:101        for line_no, raw in enumerate(handle, start=1):102            line = raw.strip()103            if not line:104                continue105            parts = line.split(None, 1)106            if len(parts) != 2:107                raise ValueError(f"invalid checksum line {line_no}: {raw!r}")108            entries.append((parts[0], parts[1]))109    return entries110 111 112def validate_checksums(package_root: Path, manifest_path: Path) -> int:113    entries = read_checksum_manifest(manifest_path)114    for expected_hash, relative_path in entries:115        target = package_root / relative_path116        require_file(target)117        if sha256_file(target) != expected_hash:118            raise ValueError(f"checksum mismatch: {relative_path}")119    return len(entries)120 121 122def main() -> int:123    parser = argparse.ArgumentParser(description=__doc__)124    parser.add_argument("--dataset-root", default="data/MP20")125    parser.add_argument("--checksum-manifest", default="metadata/sha256_manifest.txt")126    parser.add_argument("--skip-checksum", action="store_true")127    args = parser.parse_args()128 129    package_root = Path.cwd()130    dataset_root = Path(args.dataset_root)131    raw_root = dataset_root / "raw/mp_20"132    cache_root = dataset_root / "cache/mp_20"133 134    for split, expected_count in EXPECTED_SPLITS.items():135        validate_csv(raw_root / f"{split}.csv", expected_count)136        validate_cache(cache_root / split, expected_count)137 138    checksum_count = 0139    if not args.skip_checksum:140        checksum_count = validate_checksums(package_root, Path(args.checksum_manifest))141 142    print("MP-20 dataset validation passed")143    print(f"structures: {sum(EXPECTED_SPLITS.values())}")144    print(f"splits: {EXPECTED_SPLITS}")145    print(f"checksum entries verified: {checksum_count}")146    return 0147 148 149if __name__ == "__main__":150    try:151        raise SystemExit(main())152    except Exception as exc:153        print(f"MP-20 dataset validation failed: {exc}", file=sys.stderr)154        raise SystemExit(1)155