55from collections import Counter
66from dftio .constants import orbitalId , ABACUS2DFTIO
77import ase
8+ from ase .io import read
89import dpdata
910import os
1011import numpy as np
1112from dftio .io .parse import Parser , ParserRegister , find_target_line
1213from dftio .data import _keys
1314from dftio .register import Register
15+ import lmdb
16+ import pickle
17+ import shutil
18+
1419
1520@ParserRegister .register ("abacus" )
1621class AbacusParser (Parser ):
@@ -21,8 +26,9 @@ def __init__(
2126 ** kwargs
2227 ):
2328 super (AbacusParser , self ).__init__ (root , prefix )
24- if self .get_mode (idx = 0 ) == 'nscf' :
25- self .raw_sys = [dpdata .System (self .raw_datas [idx ]+ '/STRU' , fmt = 'abacus/stru' ) for idx in range (len (self .raw_datas ))]
29+ mode = self .get_mode (idx = 0 )
30+ if mode in ['nscf' , "scf" ]:
31+ 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 ))]
2632 else :
2733 self .raw_sys = [dpdata .LabeledSystem (self .raw_datas [idx ], fmt = 'abacus/' + self .get_mode (idx )) for idx in range (len (self .raw_datas ))]
2834
@@ -190,7 +196,7 @@ def get_blocks(self, idx, hamiltonian=True, overlap=False, density_matrix=False)
190196 else :
191197 raise ValueError (f'{ line } is not supported' )
192198
193- if mode == "scf" :
199+ if mode in [ "scf" , "nscf" ] :
194200 if hamiltonian :
195201 hamiltonian_dict , tmp = self .parse_matrix (
196202 matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "data-HR-sparse_SPIN0.csr" ),
@@ -242,30 +248,54 @@ def get_blocks(self, idx, hamiltonian=True, overlap=False, density_matrix=False)
242248 if hamiltonian :
243249 hamiltonian_dict = []
244250 for i in range (sys .get_nframes ()):
245- hamil , tmp = self .parse_matrix (
246- matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" + str (i )+ "_data-HR-sparse_SPIN0.csr" ),
247- nsites = nsites ,
248- site_norbits = site_norbits ,
249- orbital_types_dict = orbital_types_dict ,
250- element = element ,
251- factor = 13.605698 , # Ryd2eV
252- spinful = spinful
253- )
251+ if os .path .exists (os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" )):
252+ hamil , tmp = self .parse_matrix (
253+ matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" + str (i )+ "_data-HR-sparse_SPIN0.csr" ),
254+ nsites = nsites ,
255+ site_norbits = site_norbits ,
256+ orbital_types_dict = orbital_types_dict ,
257+ element = element ,
258+ factor = 13.605698 , # Ryd2eV
259+ spinful = spinful
260+ )
261+ else :
262+ hamil , tmp = self .parse_matrix (
263+ matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS/data-HR-sparse_SPIN0.csr" ),
264+ nsites = nsites ,
265+ site_norbits = site_norbits ,
266+ orbital_types_dict = orbital_types_dict ,
267+ element = element ,
268+ factor = 13.605698 , # Ryd2eV
269+ spinful = spinful ,
270+ step = i
271+ )
254272 assert tmp == int (np .sum (site_norbits )) * (1 + spinful )
255273 hamiltonian_dict .append (hamil )
256274
257275 if overlap :
258276 overlap_dict = []
259277 for i in range (sys .get_nframes ()):
260- ovp , tmp = self .parse_matrix (
261- matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" + str (i )+ "_data-SR-sparse_SPIN0.csr" ),
262- nsites = nsites ,
263- site_norbits = site_norbits ,
264- orbital_types_dict = orbital_types_dict ,
265- element = element ,
266- factor = 1 ,
267- spinful = spinful
268- )
278+ if os .path .exists (os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" )):
279+ ovp , tmp = self .parse_matrix (
280+ matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" + str (i )+ "_data-SR-sparse_SPIN0.csr" ),
281+ nsites = nsites ,
282+ site_norbits = site_norbits ,
283+ orbital_types_dict = orbital_types_dict ,
284+ element = element ,
285+ factor = 1 ,
286+ spinful = spinful
287+ )
288+ else :
289+ ovp , tmp = self .parse_matrix (
290+ matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS/data-SR-sparse_SPIN0.csr" ),
291+ nsites = nsites ,
292+ site_norbits = site_norbits ,
293+ orbital_types_dict = orbital_types_dict ,
294+ element = element ,
295+ factor = 1 ,
296+ spinful = spinful ,
297+ step = i
298+ )
269299 assert tmp == int (np .sum (site_norbits )) * (1 + spinful )
270300
271301 if spinful :
@@ -279,36 +309,63 @@ def get_blocks(self, idx, hamiltonian=True, overlap=False, density_matrix=False)
279309 if density_matrix :
280310 density_matrix_dict = []
281311 for i in range (sys .get_nframes ()):
282- dm , tmp = self .parse_matrix (
283- matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" + str (i )+ "_data-DMR-sparse_SPIN0.csr" ),
284- nsites = nsites ,
285- site_norbits = site_norbits ,
286- orbital_types_dict = orbital_types_dict ,
287- element = element ,
288- factor = 1 ,
289- spinful = spinful
290- )
312+ if os .path .exists (os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" )):
313+ dm , tmp = self .parse_matrix (
314+ matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS" , "matrix/" + str (i )+ "_data-DMR-sparse_SPIN0.csr" ),
315+ nsites = nsites ,
316+ site_norbits = site_norbits ,
317+ orbital_types_dict = orbital_types_dict ,
318+ element = element ,
319+ factor = 1 ,
320+ spinful = spinful
321+ )
322+ else :
323+ dm , tmp = self .parse_matrix (
324+ matrix_path = os .path .join (self .raw_datas [idx ], "OUT.ABACUS/data-DMR-sparse_SPIN0.csr" ),
325+ nsites = nsites ,
326+ site_norbits = site_norbits ,
327+ orbital_types_dict = orbital_types_dict ,
328+ element = element ,
329+ factor = 1 ,
330+ spinful = spinful ,
331+ step = i
332+ )
333+
291334 assert tmp == int (np .sum (site_norbits )) * (1 + spinful )
292335 density_matrix_dict .append (dm )
293336 else :
294337 raise NotImplementedError ("mode {} is not supported." .format (mode ))
295338
296339 return hamiltonian_dict , overlap_dict , density_matrix_dict
297340
298- def parse_matrix (self , matrix_path , nsites , site_norbits , orbital_types_dict , element , factor , spinful = False ):
341+ def parse_matrix (self , matrix_path , nsites , site_norbits , orbital_types_dict , element , factor , spinful = False , step = 0 ):
299342 site_norbits_cumsum = np .cumsum (site_norbits )
300343 norbits = int (np .sum (site_norbits ))
301344 matrix_dict = dict ()
302345 with open (matrix_path , 'r' ) as f :
303346 line = f .readline () # read "Matrix Dimension of ..."
304347 if not "Matrix Dimension of" in line :
348+ """In this case, the starting of the file is STEP 0"""
349+ # find the correct step
350+ step_found = False
351+ while line and not step_found :
352+ if "STEP" in line :
353+ stp = int (line .split ()[- 1 ])
354+ if stp != step :
355+ line = f .readline ()
356+ else :
357+ step_found = True
358+ else :
359+ line = f .readline ()
305360 line = f .readline () # ABACUS >= 3.0
306361 assert "Matrix Dimension of" in line
362+ else :
363+ assert step == 0
307364 f .readline () # read "Matrix number of ..."
308365 norbits = int (line .split ()[- 1 ])
309366 for line in f :
310367 line1 = line .split ()
311- if len (line1 ) == 0 :
368+ if len (line1 ) == 0 or len ( line1 ) == 2 :
312369 break
313370 num_element = int (line1 [3 ])
314371 if num_element != 0 :
@@ -357,4 +414,82 @@ def transform(self, mat, l_lefts, l_rights):
357414
358415 return block_lefts @ mat @ block_rights .T
359416
360-
417+ def get_abs_h0_folders (self , h0_root ):
418+ # Build a map of all directory names to their full paths to avoid repeated os.walk calls
419+ folder_path_map = {}
420+ for sub_root , dirs , _ in os .walk (h0_root ):
421+ for dirname in dirs :
422+ folder_path_map [dirname ] = os .path .join (sub_root , dirname )
423+
424+ abs_h0_folders = []
425+ valid_idx_list = []
426+ for valid_idx , a_H_folder in enumerate (self .raw_datas ):
427+ a_leaf_folder_name = os .path .split (a_H_folder )[- 1 ]
428+
429+ # Use pre-built map for O(1) lookup instead of calling slow find_leaf_folder()
430+ a_leaf_folder_path = folder_path_map .get (a_leaf_folder_name )
431+ found_flag = a_leaf_folder_path is not None
432+
433+ if found_flag :
434+ abs_h0_folders .append (a_leaf_folder_path )
435+ valid_idx_list .append (valid_idx )
436+ else :
437+ abs_h0_folders .append (None )
438+ valid_idx_list .append (None )
439+ return abs_h0_folders , valid_idx_list
440+
441+ def add_h0_delta_h (self , h0_src_root , old_lmdb_path , new_lmdb_path , keep_old_lmdb : bool = True ,
442+ keep_delta_ham_only : bool = True ):
443+ h0_root = os .path .abspath (h0_src_root )
444+ self .raw_datas , valid_idx_list = self .get_abs_h0_folders (h0_root = h0_root )
445+
446+ os .makedirs (new_lmdb_path , exist_ok = True )
447+ counter = 0
448+ batch_size = 50
449+
450+ # Process in batches
451+ for batch_start in tqdm (range (0 , len (valid_idx_list ), batch_size ), desc = "Processing batches" ):
452+ batch_end = min (batch_start + batch_size , len (valid_idx_list ))
453+ batch_indices = valid_idx_list [batch_start :batch_end ]
454+
455+ # Open LMDB environments for each batch
456+ old_db_env = lmdb .open (old_lmdb_path , readonly = True , lock = False )
457+ new_db_env = lmdb .open (new_lmdb_path , map_size = 1048576000000 , lock = True )
458+
459+ with old_db_env .begin () as old_txn , new_db_env .begin (write = True ) as new_txn :
460+ for idx in batch_indices :
461+ if idx == None :
462+ continue
463+ # Get data from old LMDB
464+ data_dict = old_txn .get (idx .to_bytes (length = 4 , byteorder = 'big' ))
465+ data_dict = pickle .loads (data_dict )
466+
467+ # Get H0 block
468+ h0_block , _0 , _ = self .get_blocks (idx , hamiltonian = True , overlap = False , density_matrix = False )
469+ h0_block = h0_block [0 ]
470+ old_ham_block = data_dict ['hamiltonian' ]
471+
472+ # Calculate delta blocks
473+ delta_block = dict ()
474+ for a_block_name in old_ham_block .keys ():
475+ a_delta_block = old_ham_block [a_block_name ] - h0_block [a_block_name ]
476+ delta_block [a_block_name ] = a_delta_block
477+
478+ # Update data dictionary
479+ data_dict ['hamiltonian' ] = delta_block
480+ if not keep_delta_ham_only :
481+ data_dict ['hamiltonian_full' ] = old_ham_block
482+ data_dict ['hamiltonian_0' ] = h0_block
483+
484+ # Store in new LMDB
485+ data_dict = pickle .dumps (data_dict )
486+ new_txn .put (counter .to_bytes (length = 4 , byteorder = 'big' ), data_dict )
487+ counter = counter + 1
488+
489+ # Close LMDB environments after each batch
490+ old_db_env .close ()
491+ new_db_env .close ()
492+
493+ # Remove old LMDB if requested
494+ if not keep_old_lmdb :
495+ shutil .rmtree (old_lmdb_path )
0 commit comments