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
39 changes: 30 additions & 9 deletions dpgen/collect/collect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -33,17 +48,22 @@ 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")
)
# 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]"))
Expand Down Expand Up @@ -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))
Expand Down
106 changes: 22 additions & 84 deletions dpgen/tools/collect_data.py
Original file line number Diff line number Diff line change
@@ -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(
Comment thread
njzjz-bot marked this conversation as resolved.
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(
Expand Down
60 changes: 60 additions & 0 deletions tests/test_collect.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
28 changes: 28 additions & 0 deletions tests/tools/test_collect_data.py
Original file line number Diff line number Diff line change
@@ -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()