Skip to content

Commit d2eb6df

Browse files
author
Han Wang
committed
support standard dpdata multisystems API
1 parent 9ae65de commit d2eb6df

3 files changed

Lines changed: 147 additions & 164 deletions

File tree

dpdata/lmdb/format.py

Lines changed: 133 additions & 145 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,6 @@
77
import msgpack_numpy as m
88
import numpy as np
99

10-
import dpdata
1110
from dpdata.format import Format
1211

1312
m.patch()
@@ -35,17 +34,8 @@ class LMDBFormat(Format):
3534
systems (with potentially different numbers of atoms) are stored in a
3635
single LMDB database file.
3736
38-
For single systems, the standard ``dpdata.System.to('lmdb', ...)`` and
39-
``dpdata.System('...', fmt='lmdb')`` APIs can be used.
40-
41-
.. note::
42-
43-
The standard ``dpdata.MultiSystems.to()`` and ``dpdata.MultiSystems()``
44-
constructor are not supported for this format. This is due to an
45-
architectural limitation, as those APIs are designed for a "directory
46-
of files" paradigm, whereas this format creates a single, unified
47-
database file. To save and load multiple systems, you must call the
48-
format's methods directly, as shown in the examples below.
37+
Both single systems and multiple systems are supported via the standard
38+
``dpdata`` APIs.
4939
5040
Examples
5141
--------
@@ -61,125 +51,146 @@ class LMDBFormat(Format):
6151
6252
**Saving multiple systems to a single LMDB database**
6353
64-
>>> from dpdata.plugins.lmdb import LMDBFormat
54+
>>> import dpdata
6555
>>> system_1 = dpdata.LabeledSystem("path/to/system1/OUTCAR", fmt="vasp/outcar")
6656
>>> system_2 = dpdata.LabeledSystem("path/to/system2/OUTCAR", fmt="vasp/outcar")
6757
>>> multi_systems_obj = dpdata.MultiSystems(system_1, system_2)
68-
>>> lmdb_formatter = LMDBFormat()
69-
>>> lmdb_formatter.to_multi_systems(
70-
... list(multi_systems_obj.systems.values()), "my_multi_system_db.lmdb"
71-
... )
58+
>>> multi_systems_obj.to("lmdb", "my_multi_system_db.lmdb")
7259
7360
**Loading multiple systems from a single LMDB database**
7461
75-
>>> from dpdata.plugins.lmdb import LMDBFormat
76-
>>> lmdb_formatter = LMDBFormat()
77-
>>> loaded_multi_systems = lmdb_formatter.from_multi_systems("my_multi_system_db.lmdb")
62+
>>> import dpdata
63+
>>> loaded_multi_systems = dpdata.MultiSystems.from_file("my_multi_system_db.lmdb", fmt="lmdb")
7864
"""
7965

8066
def to_multi_systems(
81-
self, systems, file_name, map_size=1000000000, frame_idx_fmt="012d", **kwargs
67+
self, formulas, directory, map_size=1000000000, frame_idx_fmt="012d", **kwargs
8268
):
83-
"""Save multiple systems to a single LMDB database.
69+
"""Implement MultiSystems.to for LMDB format.
8470
8571
Parameters
8672
----------
87-
systems : list of dpdata.System
88-
A list of System objects to be saved.
89-
file_name : str
90-
The path to the LMDB database directory. It will be created if it
91-
doesn't exist.
73+
formulas : list[str]
74+
list of formulas
75+
directory : str
76+
directory of system
9277
map_size : int, optional
9378
Maximum size of the LMDB database in bytes. Default is 1GB.
9479
frame_idx_fmt : str, optional
9580
The format string used to encode the frame index as a key. Default is "012d".
96-
"""
97-
from dpdata.data_type import Axis
98-
99-
os.makedirs(file_name, exist_ok=True)
100-
with lmdb.open(file_name, map_size=map_size) as env:
101-
global_frame_idx = 0
102-
system_info = []
81+
**kwargs : dict
82+
other parameters
10383
84+
Yields
85+
------
86+
tuple
87+
(self, formula) to be used by to_system
88+
"""
89+
self._frame_idx_fmt = frame_idx_fmt
90+
self._global_frame_idx = 0
91+
self._system_info = []
92+
os.makedirs(directory, exist_ok=True)
93+
with lmdb.open(directory, map_size=map_size) as env:
10494
with env.begin(write=True) as txn:
105-
for system_obj in systems:
106-
data = system_obj.data
107-
nframes = system_obj.get_nframes()
108-
formula = system_obj.formula
109-
110-
# Identify symbolic shapes and frame-dependent keys
111-
data_shapes = {}
112-
frame_dependent_keys = []
113-
for dt in type(system_obj).DTYPES:
114-
if dt.name in data:
115-
if dt.shape is not None:
116-
data_shapes[dt.name] = [
117-
s.value if isinstance(s, Axis) else s
118-
for s in dt.shape
119-
]
120-
if Axis.NFRAMES in dt.shape:
121-
frame_dependent_keys.append(dt.name)
122-
else:
123-
data_shapes[dt.name] = None
124-
125-
system_info.append(
126-
{
127-
"formula": formula,
128-
"natoms": system_obj.get_atom_numbs(),
129-
"nframes": nframes,
130-
"start_idx": global_frame_idx,
131-
"data_shapes": data_shapes,
132-
"frame_dependent_keys": frame_dependent_keys,
133-
}
134-
)
135-
136-
for i in range(nframes):
137-
frame_data = {}
138-
for key, val in data.items():
139-
if key in frame_dependent_keys:
140-
frame_data[key] = val[i]
141-
else:
142-
frame_data[key] = val
143-
144-
key = f"{global_frame_idx:{frame_idx_fmt}}".encode("ascii")
145-
value = msgpack.packb(frame_data, use_bin_type=True)
146-
txn.put(key, value)
147-
global_frame_idx += 1
148-
95+
self._txn = txn
96+
for ff in formulas:
97+
yield (self, ff)
98+
# Finalize metadata
14999
metadata = {
150-
"nframes": global_frame_idx,
151-
"system_info": system_info,
152-
"frame_idx_fmt": frame_idx_fmt,
100+
"nframes": self._global_frame_idx,
101+
"system_info": self._system_info,
102+
"frame_idx_fmt": self._frame_idx_fmt,
153103
}
154104
txn.put(b"__metadata__", msgpack.packb(metadata, use_bin_type=True))
105+
self._txn = None
155106

156-
def to_labeled_system(self, data, file_name, **kwargs):
157-
"""Save a single LabeledSystem to an LMDB database.
107+
def _dump_to_txn(self, data, txn, formula, dtypes):
108+
from dpdata.data_type import Axis
158109

159-
Parameters
160-
----------
161-
data : dict
162-
The data dictionary of a LabeledSystem.
163-
file_name : str
164-
The path to the LMDB database directory.
165-
"""
110+
nframes = data["coords"].shape[0]
111+
112+
# Identify symbolic shapes and frame-dependent keys
113+
data_shapes = {}
114+
frame_dependent_keys = []
115+
for dt in dtypes:
116+
if dt.name in data:
117+
if dt.shape is not None:
118+
data_shapes[dt.name] = [
119+
s.value if isinstance(s, Axis) else s for s in dt.shape
120+
]
121+
if Axis.NFRAMES in dt.shape:
122+
frame_dependent_keys.append(dt.name)
123+
else:
124+
data_shapes[dt.name] = None
125+
126+
# Record system info
127+
# natoms needs to be extracted from data
128+
if "atom_numbs" in data:
129+
natoms_list = data["atom_numbs"]
130+
else:
131+
# Fallback for systems without atom_numbs (should not happen in valid dpdata systems)
132+
natoms_list = []
133+
134+
self._system_info.append(
135+
{
136+
"formula": formula,
137+
"natoms": natoms_list,
138+
"nframes": nframes,
139+
"start_idx": self._global_frame_idx,
140+
"data_shapes": data_shapes,
141+
"frame_dependent_keys": frame_dependent_keys,
142+
}
143+
)
144+
145+
for i in range(nframes):
146+
frame_data = {}
147+
for key, val in data.items():
148+
if key in frame_dependent_keys:
149+
frame_data[key] = val[i]
150+
else:
151+
frame_data[key] = val
152+
153+
key = f"{self._global_frame_idx:{self._frame_idx_fmt}}".encode("ascii")
154+
value = msgpack.packb(frame_data, use_bin_type=True)
155+
txn.put(key, value)
156+
self._global_frame_idx += 1
157+
158+
def to_labeled_system(self, data, file_name, **kwargs):
159+
"""Save a single LabeledSystem to an LMDB database."""
166160
from dpdata.system import LabeledSystem
167161

168-
self.to_multi_systems([LabeledSystem(data=data)], file_name, **kwargs)
162+
if isinstance(file_name, tuple) and file_name[0] is self:
163+
txn, formula = self._txn, file_name[1]
164+
self._dump_to_txn(data, txn, formula, LabeledSystem.DTYPES)
165+
else:
166+
# Single system call: use to_multi_systems logic
167+
# Infer formula from data if possible, or use default
168+
formula = kwargs.get("formula", "unknown")
169+
gen = self.to_multi_systems([formula], file_name, **kwargs)
170+
handle = next(gen)
171+
self.to_labeled_system(data, handle, **kwargs)
172+
try:
173+
next(gen)
174+
except StopIteration:
175+
pass
169176

170177
def to_system(self, data, file_name, **kwargs):
171-
"""Save a single System to an LMDB database.
172-
173-
Parameters
174-
----------
175-
data : dict
176-
The data dictionary of a System.
177-
file_name : str
178-
The path to the LMDB database directory.
179-
"""
178+
"""Save a single System to an LMDB database."""
180179
from dpdata.system import System
181180

182-
self.to_multi_systems([System(data=data)], file_name, **kwargs)
181+
if isinstance(file_name, tuple) and file_name[0] is self:
182+
txn, formula = self._txn, file_name[1]
183+
self._dump_to_txn(data, txn, formula, System.DTYPES)
184+
else:
185+
# Single system call
186+
formula = kwargs.get("formula", "unknown")
187+
gen = self.to_multi_systems([formula], file_name, **kwargs)
188+
handle = next(gen)
189+
self.to_system(data, handle, **kwargs)
190+
try:
191+
next(gen)
192+
except StopIteration:
193+
pass
183194

184195
def from_multi_systems(self, file_name, map_size=1000000000, **kwargs):
185196
"""Load multiple systems from a single LMDB database.
@@ -189,19 +200,18 @@ def from_multi_systems(self, file_name, map_size=1000000000, **kwargs):
189200
file_name : str
190201
The path to the LMDB database directory.
191202
map_size : int, optional
192-
Maximum size of the LMDB database in bytes. This parameter is included
193-
for consistency with `to_multi_systems` but is generally ignored
194-
when opening in `readonly=True` mode.
195-
196-
Returns
197-
-------
198-
dpdata.MultiSystems
199-
A MultiSystems object containing all systems stored in the LMDB.
203+
Maximum size of the LMDB database in bytes.
204+
**kwargs : dict
205+
other parameters
206+
207+
Yields
208+
------
209+
dict
210+
data dictionary for each system
200211
"""
201212
from dpdata.data_type import Axis, DataType
202213
from dpdata.system import LabeledSystem, System
203214

204-
systems = []
205215
with lmdb.open(file_name, readonly=True) as env:
206216
with env.begin() as txn:
207217
metadata_packed = txn.get(b"__metadata__")
@@ -257,42 +267,20 @@ def from_multi_systems(self, file_name, map_size=1000000000, **kwargs):
257267
else:
258268
agg_data[key] = val
259269

260-
systems.append(cls(data=agg_data))
261-
262-
return dpdata.MultiSystems(*systems)
270+
yield agg_data
263271

264272
def from_labeled_system(self, file_name, **kwargs):
265-
"""Load data for a single LabeledSystem from an LMDB database.
266-
267-
Parameters
268-
----------
269-
file_name : str
270-
The path to the LMDB database directory.
271-
272-
Returns
273-
-------
274-
dict
275-
The data dictionary for the loaded LabeledSystem.
276-
"""
277-
# from_multi_systems returns a MultiSystems object
278-
multisystems_obj = self.from_multi_systems(file_name, **kwargs)
279-
# We need the data dictionary of the first (and only) system for from_labeled_system
280-
return multisystems_obj[0].data
273+
"""Load data for a single LabeledSystem from an LMDB database."""
274+
if isinstance(file_name, dict):
275+
return file_name
276+
# from_multi_systems returns a generator of dicts
277+
gen = self.from_multi_systems(file_name, **kwargs)
278+
return next(gen)
281279

282280
def from_system(self, file_name, **kwargs):
283-
"""Load data for a single System from an LMDB database.
284-
285-
Parameters
286-
----------
287-
file_name : str
288-
The path to the LMDB database directory.
289-
290-
Returns
291-
-------
292-
dict
293-
The data dictionary for the loaded System.
294-
"""
295-
# from_multi_systems returns a MultiSystems object
296-
multisystems_obj = self.from_multi_systems(file_name, **kwargs)
297-
# We need the data dictionary of the first (and only) system for from_system
298-
return multisystems_obj[0].data
281+
"""Load data for a single System from an LMDB database."""
282+
if isinstance(file_name, dict):
283+
return file_name
284+
# from_multi_systems returns a generator of dicts
285+
gen = self.from_multi_systems(file_name, **kwargs)
286+
return next(gen)

0 commit comments

Comments
 (0)