Skip to content

Commit 9f89dd4

Browse files
Merge pull request #13 from Franklalalala/main
fix: add nscf density matrix support for ABACUS
2 parents c1646bd + 298290a commit 9f89dd4

2 files changed

Lines changed: 128 additions & 6 deletions

File tree

dftio/io/abacus/abacus_parser.py

Lines changed: 86 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,10 @@
1212
from dftio.io.parse import Parser, ParserRegister, find_target_line
1313
from dftio.data import _keys
1414
from dftio.register import Register
15+
import lmdb
16+
import pickle
17+
import shutil
18+
1519

1620
@ParserRegister.register("abacus")
1721
class AbacusParser(Parser):
@@ -23,9 +27,7 @@ def __init__(
2327
):
2428
super(AbacusParser, self).__init__(root, prefix)
2529
mode = self.get_mode(idx=0)
26-
if mode == 'nscf':
27-
self.raw_sys = [dpdata.System(self.raw_datas[idx]+'/STRU', fmt='abacus/stru') for idx in range(len(self.raw_datas))]
28-
elif mode == "scf":
30+
if mode in ['nscf', "scf"]:
2931
self.raw_sys = [dpdata.System(read(os.path.join(self.raw_datas[idx], "OUT.ABACUS", "STRU.cif")), fmt="ase/structure") for idx in range(len(self.raw_datas))]
3032
else:
3133
self.raw_sys = [dpdata.LabeledSystem(self.raw_datas[idx], fmt='abacus/'+self.get_mode(idx)) for idx in range(len(self.raw_datas))]
@@ -194,7 +196,7 @@ def get_blocks(self, idx, hamiltonian=True, overlap=False, density_matrix=False)
194196
else:
195197
raise ValueError(f'{line} is not supported')
196198

197-
if mode == "scf":
199+
if mode in ["scf", "nscf"]:
198200
if hamiltonian:
199201
hamiltonian_dict, tmp = self.parse_matrix(
200202
matrix_path=os.path.join(self.raw_datas[idx], "OUT.ABACUS", "data-HR-sparse_SPIN0.csr"),
@@ -312,7 +314,7 @@ def parse_matrix(self, matrix_path, nsites, site_norbits, orbital_types_dict, el
312314
norbits = int(line.split()[-1])
313315
for line in f:
314316
line1 = line.split()
315-
if len(line1) == 0:
317+
if len(line1) == 0 or len(line1) == 2:
316318
break
317319
num_element = int(line1[3])
318320
if num_element != 0:
@@ -361,4 +363,82 @@ def transform(self, mat, l_lefts, l_rights):
361363

362364
return block_lefts @ mat @ block_rights.T
363365

364-
366+
def get_abs_h0_folders(self, h0_root):
367+
# Build a map of all directory names to their full paths to avoid repeated os.walk calls
368+
folder_path_map = {}
369+
for sub_root, dirs, _ in os.walk(h0_root):
370+
for dirname in dirs:
371+
folder_path_map[dirname] = os.path.join(sub_root, dirname)
372+
373+
abs_h0_folders = []
374+
valid_idx_list = []
375+
for valid_idx, a_H_folder in enumerate(self.raw_datas):
376+
a_leaf_folder_name = os.path.split(a_H_folder)[-1]
377+
378+
# Use pre-built map for O(1) lookup instead of calling slow find_leaf_folder()
379+
a_leaf_folder_path = folder_path_map.get(a_leaf_folder_name)
380+
found_flag = a_leaf_folder_path is not None
381+
382+
if found_flag:
383+
abs_h0_folders.append(a_leaf_folder_path)
384+
valid_idx_list.append(valid_idx)
385+
else:
386+
abs_h0_folders.append(None)
387+
valid_idx_list.append(None)
388+
return abs_h0_folders, valid_idx_list
389+
390+
def add_h0_delta_h(self, h0_src_root, old_lmdb_path, new_lmdb_path, keep_old_lmdb: bool = True,
391+
keep_delta_ham_only: bool = True):
392+
h0_root = os.path.abspath(h0_src_root)
393+
self.raw_datas, valid_idx_list = self.get_abs_h0_folders(h0_root=h0_root)
394+
395+
os.makedirs(new_lmdb_path, exist_ok=True)
396+
counter = 0
397+
batch_size = 50
398+
399+
# Process in batches
400+
for batch_start in tqdm(range(0, len(valid_idx_list), batch_size), desc="Processing batches"):
401+
batch_end = min(batch_start + batch_size, len(valid_idx_list))
402+
batch_indices = valid_idx_list[batch_start:batch_end]
403+
404+
# Open LMDB environments for each batch
405+
old_db_env = lmdb.open(old_lmdb_path, readonly=True, lock=False)
406+
new_db_env = lmdb.open(new_lmdb_path, map_size=1048576000000, lock=True)
407+
408+
with old_db_env.begin() as old_txn, new_db_env.begin(write=True) as new_txn:
409+
for idx in batch_indices:
410+
if idx == None:
411+
continue
412+
# Get data from old LMDB
413+
data_dict = old_txn.get(idx.to_bytes(length=4, byteorder='big'))
414+
data_dict = pickle.loads(data_dict)
415+
416+
# Get H0 block
417+
h0_block, _0, _ = self.get_blocks(idx, hamiltonian=True, overlap=False, density_matrix=False)
418+
h0_block = h0_block[0]
419+
old_ham_block = data_dict['hamiltonian']
420+
421+
# Calculate delta blocks
422+
delta_block = dict()
423+
for a_block_name in old_ham_block.keys():
424+
a_delta_block = old_ham_block[a_block_name] - h0_block[a_block_name]
425+
delta_block[a_block_name] = a_delta_block
426+
427+
# Update data dictionary
428+
data_dict['hamiltonian'] = delta_block
429+
if not keep_delta_ham_only:
430+
data_dict['hamiltonian_full'] = old_ham_block
431+
data_dict['hamiltonian_0'] = h0_block
432+
433+
# Store in new LMDB
434+
data_dict = pickle.dumps(data_dict)
435+
new_txn.put(counter.to_bytes(length=4, byteorder='big'), data_dict)
436+
counter = counter + 1
437+
438+
# Close LMDB environments after each batch
439+
old_db_env.close()
440+
new_db_env.close()
441+
442+
# Remove old LMDB if requested
443+
if not keep_old_lmdb:
444+
shutil.rmtree(old_lmdb_path)

test/test_abacus_add_h0.py

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,42 @@
1+
import os
2+
import shutil
3+
from dftio.io.abacus.abacus_parser import AbacusParser
4+
from tqdm import tqdm
5+
6+
7+
def main():
8+
parser = AbacusParser(
9+
root=r'/share/mp_20_abacus_production/15_elements_0524/train_50_nscf/cooking',
10+
prefix=r'*/cooking/db_seq_id_*',
11+
)
12+
num_entries = len(parser.raw_datas)
13+
outroot = r'/share/mp_20_abacus_production/15_elements_0524/50_H0_workspace_v2/raw_lmdb_data'
14+
15+
if os.path.exists(outroot):
16+
import shutil
17+
shutil.rmtree(outroot)
18+
19+
raw_outroot = os.path.join(outroot, 'raw')
20+
old_lmdb_path = os.path.join(raw_outroot, "data.{}.lmdb".format(os.getpid()))
21+
new_lmdb_path = os.path.join(outroot, 'h0', "data.h0.lmdb")
22+
for idx in tqdm(range(num_entries)):
23+
parser.write(
24+
idx=int(idx),
25+
format='lmdb',
26+
hamiltonian=True,
27+
overlap=True,
28+
outroot=raw_outroot,
29+
eigenvalue=True,
30+
density_matrix=True,
31+
band_index_min=0
32+
)
33+
parser.add_h0_delta_h(
34+
h0_src_root= r'/share/mp_20_abacus_production/15_elements_0524/50_H0_workspace_v2/cooking',
35+
old_lmdb_path=old_lmdb_path,
36+
new_lmdb_path=new_lmdb_path,
37+
keep_old_lmdb=True
38+
)
39+
40+
41+
if __name__ == "__main__":
42+
main()

0 commit comments

Comments
 (0)