From a44808b2cda203e1336975bccde6c48027b97d8d Mon Sep 17 00:00:00 2001 From: "weihong.xu" Date: Wed, 24 Dec 2025 10:44:09 +0800 Subject: [PATCH] support specific files by glob --- mace/cli/preprocess_data.py | 8 ++++---- mace/tools/utils.py | 19 ++++++++++++++++++- 2 files changed, 22 insertions(+), 5 deletions(-) diff --git a/mace/cli/preprocess_data.py b/mace/cli/preprocess_data.py index 84dad4944..a832c64c2 100644 --- a/mace/cli/preprocess_data.py +++ b/mace/cli/preprocess_data.py @@ -22,7 +22,7 @@ from mace.modules import compute_statistics from mace.tools import torch_geometric from mace.tools.scripts_utils import get_atomic_energies, get_dataset_from_xyz -from mace.tools.utils import AtomicNumberTable +from mace.tools.utils import AtomicNumberTable, expand_glob def compute_stats_target( @@ -176,11 +176,11 @@ def run(args: argparse.Namespace): # Data preparation collections, atomic_energies_dict = get_dataset_from_xyz( work_dir=args.work_dir, - train_path=args.train_file, - valid_path=args.valid_file, + train_path=expand_glob(args.train_file), + valid_path=expand_glob(args.valid_file), valid_fraction=args.valid_fraction, config_type_weights=config_type_weights, - test_path=args.test_file, + test_path=expand_glob(args.test_file), seed=args.seed, key_specification=args.key_specification, head_name="", diff --git a/mace/tools/utils.py b/mace/tools/utils.py index ae2c9bf3f..b8a659484 100644 --- a/mace/tools/utils.py +++ b/mace/tools/utils.py @@ -9,10 +9,11 @@ import os import sys from pathlib import Path -from typing import Any, Dict, Iterable, Optional, Sequence, Union +from typing import Any, Dict, Iterable, Optional, Sequence, Union, List import numpy as np import torch +from glob import glob from .torch_tools import to_numpy @@ -205,3 +206,19 @@ def filter_nonzero_weight( quantity_l[-1] = filtered_q return 1.0 + + +def expand_glob(file_path: Optional[Union[str, List[str]]]): + if file_path is None: + return file_path + if isinstance(file_path, str): + file_path = [file_path] + if not isinstance(file_path, list): + return file_path + expanded_paths = [] + for path in file_path: + _files = glob(path) + if not _files: + raise FileNotFoundError(f"No files matched the pattern: {path}") + expanded_paths.extend(_files) + return expanded_paths