diff --git a/dpgen/collect/collect.py b/dpgen/collect/collect.py index db2a6a163..19c77cdbf 100644 --- a/dpgen/collect/collect.py +++ b/dpgen/collect/collect.py @@ -12,8 +12,23 @@ def collect_data( - target_folder, param_file, output, verbose=True, shuffle=True, merge=True + target_folder, + param_file, + output, + verbose=True, + shuffle=True, + merge=True, + include_init_data=True, + iter_output_prefix="sys.", + discover_existing_iters=False, ): + """Collect initial and iterative datasets from a DP-GEN job. + + ``include_init_data``, ``iter_output_prefix``, and + ``discover_existing_iters`` preserve the historical behavior used by + :mod:`dpgen.tools.collect_data` without duplicating the data loading and + serialization implementation. + """ target_folder = os.path.abspath(target_folder) output = os.path.abspath(output) # goto input @@ -33,7 +48,7 @@ def collect_data( # init systems init_data = [] init_data_prefix = jdata.get("init_data_prefix", "") - init_data_sys = jdata.get("init_data_sys", []) + init_data_sys = jdata.get("init_data_sys", []) if include_init_data else [] for ii in init_data_sys: init_data.append( dpdata.LabeledSystem(os.path.join(init_data_prefix, ii), fmt="deepmd/npy") @@ -41,9 +56,14 @@ def collect_data( # collect systems from iter dirs coll_data = {} numb_sys = len(sys_configs) - model_devi_jobs = jdata.get("model_devi_jobs", {}) - numb_jobs = len(model_devi_jobs) - iters = ["iter.%06d" % ii for ii in range(numb_jobs)] # noqa: UP031 + if discover_existing_iters: + # The deprecated collector operated on completed directories rather + # than assuming one iteration per current model_devi_jobs entry. + iters = sorted(glob.glob("iter.[0-9]*[0-9]")) + else: + model_devi_jobs = jdata.get("model_devi_jobs", {}) + numb_jobs = len(model_devi_jobs) + iters = ["iter.%06d" % ii for ii in range(numb_jobs)] # noqa: UP031 # loop over iters to collect data for ii in range(len(iters)): iter_data = glob.glob(os.path.join(iters[ii], "02.fp", "data.[0-9]*[0-9]")) @@ -95,12 +115,13 @@ def collect_data( os.chdir(cwd) os.makedirs(output, exist_ok=True) # dump init data - for idx, ii in enumerate(init_data): - out_dir = "init." + (data_system_fmt % idx) - ii.to("deepmd/npy", os.path.join(output, out_dir)) + if include_init_data: + for idx, ii in enumerate(init_data): + out_dir = "init." + (data_system_fmt % idx) + ii.to("deepmd/npy", os.path.join(output, out_dir)) # dump iter data for kk in coll_data.keys(): - out_dir = f"sys.{kk}" + out_dir = f"{iter_output_prefix}{kk}" nframes = coll_data[kk].get_nframes() coll_data[kk].to("deepmd/npy", os.path.join(output, out_dir), set_size=nframes) # coll_data[kk].to('deepmd/npy', os.path.join(output, out_dir)) diff --git a/dpgen/tools/collect_data.py b/dpgen/tools/collect_data.py index 9c4a18847..9b5ac1e92 100755 --- a/dpgen/tools/collect_data.py +++ b/dpgen/tools/collect_data.py @@ -1,98 +1,36 @@ #!/usr/bin/env python3 import argparse -import glob -import json -import os -import subprocess as sp +import warnings - -def file_len(fname): - with open(fname) as f: - for i, l in enumerate(f): - pass - return i + 1 +from dpgen.collect.collect import collect_data as collect_current_data def collect_data(target_folder, param_file, output, verbose=True): - target_folder = os.path.abspath(target_folder) - output = os.path.abspath(output) - tool_path = os.path.join( - os.path.dirname(os.path.realpath(__file__)), "..", "template" + """Delegate the legacy helper to the maintained collection implementation.""" + warnings.warn( + "dpgen.tools.collect_data is deprecated; use `dpgen collect` or " + "`dpgen.collect.collect.collect_data` instead.", + DeprecationWarning, + stacklevel=2, + ) + return collect_current_data( + target_folder, + param_file, + output, + verbose=verbose, + shuffle=True, + merge=False, + include_init_data=False, + iter_output_prefix="system.", + discover_existing_iters=True, ) - command_cvt_2_raw = os.path.join(tool_path, "tools.vasp", "convert2raw.py") - command_cvt_2_raw += " data.configs" - command_shuffle_raw = os.path.join(tool_path, "tools.raw", "shuffle_raw.py") - command_raw_2_set = os.path.join(tool_path, "tools.raw", "raw_to_set.sh") - # goto input - cwd = os.getcwd() - os.chdir(target_folder) - jdata = json.load(open(param_file)) - sys = jdata["sys_configs"] - if verbose: - max_str_len = max([len(str(ii)) for ii in sys]) - ptr_fmt = "%%%ds %%6d" % (max_str_len + 5) # noqa: UP031 - # collect systems from iter dirs - coll_sys = [[] for ii in sys] - numb_sys = len(sys) - iters = glob.glob("iter.[0-9]*[0-9]") - iters.sort() - for ii in iters: - iter_data = glob.glob(os.path.join(ii, "02.fp", "data.[0-9]*[0-9]")) - iter_data.sort() - for jj in iter_data: - sys_idx = int(os.path.basename(jj).split(".")[-1]) - coll_sys[sys_idx].append(jj) - # create output dir - os.makedirs(output, exist_ok=True) - # loop over systems - for idx, ii in enumerate(coll_sys): - if len(ii) == 0: - continue - # link iter data dirs - out_sys_path = os.path.join(output, "system.%03d" % idx) # noqa: UP031 - os.makedirs(out_sys_path, exist_ok=True) - cwd_ = os.getcwd() - os.chdir(out_sys_path) - for jj in ii: - in_sys_path = os.path.join(target_folder, jj) - in_iter = in_sys_path.split("/")[-3] - in_base = in_sys_path.split("/")[-1] - out_file = in_iter + "." + in_base - if os.path.exists(out_file): - os.remove(out_file) - os.symlink(in_sys_path, out_file) - # cat data.configs - data_configs = glob.glob( - os.path.join("iter.[0-9]*[0-9].data.[0-9]*[0-9]", "orig", "data.configs") - ) - data_configs.sort() - os.makedirs("orig", exist_ok=True) - with open(os.path.join("orig", "data.configs"), "w") as outfile: - for fname in data_configs: - with open(fname) as infile: - outfile.write(infile.read()) - # convert to raw - os.chdir("orig") - sp.check_call(command_cvt_2_raw, shell=True) - os.chdir("..") - # shuffle raw - sp.check_call(command_shuffle_raw + " orig " + " . > /dev/null", shell=True) - if os.path.exists("type.raw"): - os.remove("type.raw") - os.symlink(os.path.join("orig", "type.raw"), "type.raw") - # raw to sets - sp.check_call(command_raw_2_set + " > /dev/null", shell=True) - # print summary - if verbose: - ndata = file_len("box.raw") - print(ptr_fmt % (str(sys[idx]), ndata)) - # ch dir - os.chdir(cwd_) def _main(): - parser = argparse.ArgumentParser(description="Collect data from DP-GEN iterations") + parser = argparse.ArgumentParser( + description="Deprecated wrapper for `dpgen collect`" + ) parser.add_argument("JOB_DIR", type=str, help="the directory of the DP-GEN job") parser.add_argument("OUTPUT", type=str, help="the output directory of data") parser.add_argument( diff --git a/tests/test_collect.py b/tests/test_collect.py index 99979697d..e25bcb212 100644 --- a/tests/test_collect.py +++ b/tests/test_collect.py @@ -35,3 +35,63 @@ def test_collect_data(self): collect_data(inpdir, param_file.name, outdir, verbose=True) ms = dpdata.MultiSystems().from_deepmd_npy(outdir) self.assertEqual(ms.get_nframes(), self.data.get_nframes() * 3) + + def test_legacy_iterative_output_layout(self): + """Compatibility mode omits initial data and keeps system.* names.""" + with ( + tempfile.TemporaryDirectory() as inpdir, + tempfile.TemporaryDirectory() as outdir, + tempfile.NamedTemporaryFile() as param_file, + ): + self.data.to_deepmd_npy(Path(inpdir) / "iter.000000" / "02.fp" / "data.000") + init_path = Path(inpdir) / "init-data" + self.data.to_deepmd_npy(init_path) + with open(param_file.name, "w") as fp: + json.dump( + { + "sys_configs": ["sys1"], + "model_devi_jobs": [{}], + "init_data_sys": [str(init_path)], + }, + fp, + ) + + collect_data( + inpdir, + param_file.name, + outdir, + verbose=False, + merge=False, + include_init_data=False, + iter_output_prefix="system.", + ) + + self.assertTrue((Path(outdir) / "system.000").is_dir()) + self.assertFalse(any(Path(outdir).glob("init.*"))) + + def test_discovers_completed_iteration_directories(self): + """Compatibility mode includes iterations beyond the parameter list.""" + with ( + tempfile.TemporaryDirectory() as inpdir, + tempfile.TemporaryDirectory() as outdir, + tempfile.NamedTemporaryFile() as param_file, + ): + self.data.to_deepmd_npy(Path(inpdir) / "iter.000000" / "02.fp" / "data.000") + self.data.to_deepmd_npy(Path(inpdir) / "iter.000001" / "02.fp" / "data.000") + with open(param_file.name, "w") as fp: + json.dump( + {"sys_configs": ["sys1"], "model_devi_jobs": [{}]}, + fp, + ) + + collect_data( + inpdir, + param_file.name, + outdir, + verbose=False, + shuffle=False, + discover_existing_iters=True, + ) + + systems = dpdata.MultiSystems().from_deepmd_npy(outdir) + self.assertEqual(systems.get_nframes(), self.data.get_nframes() * 2) diff --git a/tests/tools/test_collect_data.py b/tests/tools/test_collect_data.py new file mode 100644 index 000000000..b8070ee8d --- /dev/null +++ b/tests/tools/test_collect_data.py @@ -0,0 +1,28 @@ +import unittest +from unittest.mock import patch + +from dpgen.tools.collect_data import collect_data + + +class TestLegacyCollectData(unittest.TestCase): + @patch("dpgen.tools.collect_data.collect_current_data", return_value="result") + def test_delegates_to_maintained_collector(self, current_collect_data): + with self.assertWarnsRegex(DeprecationWarning, "dpgen collect"): + result = collect_data("job", "param.json", "output", verbose=False) + + self.assertEqual(result, "result") + current_collect_data.assert_called_once_with( + "job", + "param.json", + "output", + verbose=False, + shuffle=True, + merge=False, + include_init_data=False, + iter_output_prefix="system.", + discover_existing_iters=True, + ) + + +if __name__ == "__main__": + unittest.main()