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
99 changes: 51 additions & 48 deletions deepmd/dpmodel/utils/lmdb_data.py
Original file line number Diff line number Diff line change
Expand Up @@ -1766,60 +1766,63 @@ def merge_lmdb(

for src_path in src_paths:
src_env = _open_lmdb(src_path)
with src_env.begin() as txn:
meta = _read_metadata(txn)
nframes, src_fmt, natoms_per_type = _parse_metadata(meta)
fallback_natoms = sum(natoms_per_type)

if first_system_info is None:
first_system_info = meta.get("system_info", {})
if first_type_map is None:
first_type_map = meta.get("type_map")

# Check for pre-computed frame_nlocs in source
src_nlocs = meta.get("frame_nlocs")
# Check for frame_system_ids in source
src_sys_ids = meta.get("frame_system_ids")

with src_env.begin() as src_txn, dst_env.begin(write=True) as dst_txn:
for i in range(nframes):
src_key = format(i, src_fmt).encode()
raw = src_txn.get(src_key)
if raw is None:
continue
dst_key = format(frame_idx, fmt).encode()
dst_txn.put(dst_key, raw)
try:
with src_env.begin() as txn:
meta = _read_metadata(txn)
nframes, src_fmt, natoms_per_type = _parse_metadata(meta)
fallback_natoms = sum(natoms_per_type)

if first_system_info is None:
first_system_info = meta.get("system_info", {})
if first_type_map is None:
first_type_map = meta.get("type_map")

# Check for pre-computed frame_nlocs in source
src_nlocs = meta.get("frame_nlocs")
# Check for frame_system_ids in source
src_sys_ids = meta.get("frame_system_ids")

with src_env.begin() as src_txn, dst_env.begin(write=True) as dst_txn:
for i in range(nframes):
src_key = format(i, src_fmt).encode()
raw = src_txn.get(src_key)
if raw is None:
continue
dst_key = format(frame_idx, fmt).encode()
dst_txn.put(dst_key, raw)

# Get nloc for this frame
if src_nlocs is not None:
frame_nlocs.append(int(src_nlocs[i]))
else:
frame_raw = msgpack.unpackb(raw, raw=False)
atype_raw = frame_raw.get("atom_types")
if isinstance(atype_raw, dict):
shape = atype_raw.get("shape") or atype_raw.get(b"shape")
if shape:
frame_nlocs.append(int(shape[0]))
# Get nloc for this frame
if src_nlocs is not None:
frame_nlocs.append(int(src_nlocs[i]))
else:
frame_raw = msgpack.unpackb(raw, raw=False)
atype_raw = frame_raw.get("atom_types")
if isinstance(atype_raw, dict):
shape = atype_raw.get("shape") or atype_raw.get(b"shape")
if shape:
frame_nlocs.append(int(shape[0]))
else:
frame_nlocs.append(fallback_natoms)
else:
frame_nlocs.append(fallback_natoms)
else:
frame_nlocs.append(fallback_natoms)

# Propagate system IDs with offset
if src_sys_ids is not None and i < len(src_sys_ids):
frame_system_ids.append(int(src_sys_ids[i]) + sys_id_offset)
else:
frame_system_ids.append(sys_id_offset)

frame_idx += 1
# Propagate system IDs with offset
if src_sys_ids is not None and i < len(src_sys_ids):
frame_system_ids.append(int(src_sys_ids[i]) + sys_id_offset)
else:
frame_system_ids.append(sys_id_offset)

# Update sys_id_offset for next source
if src_sys_ids is not None and len(src_sys_ids) > 0:
sys_id_offset += max(int(s) for s in src_sys_ids) + 1
else:
sys_id_offset += 1
frame_idx += 1

src_env.close()
# Update sys_id_offset for next source
if src_sys_ids is not None and len(src_sys_ids) > 0:
sys_id_offset += max(int(s) for s in src_sys_ids) + 1
else:
sys_id_offset += 1
finally:
# Release only the reference acquired for this merge. Other readers
# may still share the cached environment and must remain usable.
_close_lmdb(src_path)

# Write merged metadata with frame_nlocs for fast init
merged_meta = {
Expand Down
43 changes: 43 additions & 0 deletions source/tests/pt/test_lmdb_dataloader.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@
Consistency tests (dpmodel vs pt) live in source/tests/consistent/test_lmdb_data.py.
"""

import gc

import lmdb
import msgpack
import numpy as np
Expand Down Expand Up @@ -751,6 +753,47 @@ def test_merge_preserves_type_map(self, tmp_path):
env.close()
assert meta.get("type_map") == ["O", "H"]

def test_merge_keeps_cached_source_readers_open(self, tmp_path):
"""A merge must release its cache reference without closing readers."""
src = str(tmp_path / "shared_src.lmdb")
dst = str(tmp_path / "shared_dst.lmdb")
_create_lmdb_with_system_ids(
src, system_frames=[3], natoms=6, type_map=["O", "H"]
)
cache_key = str(tmp_path.joinpath("shared_src.lmdb").resolve())
existing_reader = LmdbDataReader(src, ["O", "H"])
new_reader = None
expected_coord = existing_reader[0]["coord"].copy()
cached_env, refcount_before_merge = lmdb_data._ENV_CACHE[cache_key]

try:
merge_lmdb([src], dst)

# merge_lmdb shares the process-level environment cache with readers.
# Its temporary reference must be balanced without replacing or
# invalidating the environment retained by the existing reader.
env_after_merge, refcount_after_merge = lmdb_data._ENV_CACHE[cache_key]
assert env_after_merge is cached_env
assert refcount_after_merge == refcount_before_merge
np.testing.assert_array_equal(existing_reader[0]["coord"], expected_coord)

new_reader = LmdbDataReader(src, ["O", "H"])
env_after_new_reader, refcount_after_new_reader = lmdb_data._ENV_CACHE[
cache_key
]
assert env_after_new_reader is cached_env
assert refcount_after_new_reader == refcount_before_merge + 1
np.testing.assert_array_equal(new_reader[0]["coord"], expected_coord)
finally:
# Drop both strong references so their finalizers release the
# cache entries before the cleanup assertion below.
del existing_reader, new_reader
gc.collect()
# Keep a failing regression isolated from later LMDB tests. On the
# fixed path reader finalizers normally remove this entry already.
while cache_key in lmdb_data._ENV_CACHE:
lmdb_data._close_lmdb(src)


# ============================================================
# Multitask LMDB training
Expand Down
Loading