@@ -211,13 +211,13 @@ def check_blocks(self, idx, hamiltonian: bool=False, overlap: bool=False, densit
211211
212212 return True
213213
214- def write (self , idx , outroot , format , eigenvalue , hamiltonian , overlap , density_matrix , band_index_min , ** kwargs ):
214+ def write (self , idx , outroot , format , eigenvalue , hamiltonian , overlap , density_matrix , band_index_min , energy = False , ** kwargs ):
215215 if format == "hdf5" :
216- self .write_hdf5 (idx = idx , outroot = outroot , eigenvalue = eigenvalue , hamiltonian = hamiltonian , overlap = overlap , density_matrix = density_matrix ,band_index_min = band_index_min )
216+ self .write_hdf5 (idx = idx , outroot = outroot , eigenvalue = eigenvalue , hamiltonian = hamiltonian , overlap = overlap , density_matrix = density_matrix ,band_index_min = band_index_min , energy = energy )
217217 elif format in ["dat" , "ase" ]:
218- self .write_dat (idx = idx , outroot = outroot , fmt = format , eigenvalue = eigenvalue , hamiltonian = hamiltonian , overlap = overlap , density_matrix = density_matrix ,band_index_min = band_index_min )
218+ self .write_dat (idx = idx , outroot = outroot , fmt = format , eigenvalue = eigenvalue , hamiltonian = hamiltonian , overlap = overlap , density_matrix = density_matrix ,band_index_min = band_index_min , energy = energy )
219219 elif format == "lmdb" :
220- self .write_lmdb (idx = idx , outroot = outroot , eigenvalue = eigenvalue , hamiltonian = hamiltonian , overlap = overlap , density_matrix = density_matrix ,band_index_min = band_index_min )
220+ self .write_lmdb (idx = idx , outroot = outroot , eigenvalue = eigenvalue , hamiltonian = hamiltonian , overlap = overlap , density_matrix = density_matrix ,band_index_min = band_index_min , energy = energy )
221221 else :
222222 raise NotImplementedError (f"Format: { format } is not implemented!" )
223223
@@ -242,10 +242,10 @@ def write_struct(self, structure, out_dir, fmt='dat'):
242242 else :
243243 raise NotImplementedError (f"Format: { fmt } is not implemented!" )
244244
245- def write_dat (self , idx , outroot , fmt = 'dat' , eigenvalue = False , hamiltonian = False , overlap = False , density_matrix = False , band_index_min = 0 ):
245+ def write_dat (self , idx , outroot , fmt = 'dat' , eigenvalue = False , hamiltonian = False , overlap = False , density_matrix = False , band_index_min = 0 , energy = False ):
246246 # write structure
247247 os .makedirs (outroot , exist_ok = True )
248-
248+
249249 structure = self .get_structure (idx )
250250
251251 out_dir = os .path .join (outroot , self .formula (idx = idx )+ ".{}" .format (idx ))
@@ -255,7 +255,7 @@ def write_dat(self, idx, outroot, fmt='dat', eigenvalue=False, hamiltonian=False
255255 # np.savetxt(os.path.join(out_dir, "positions.dat"), structure[_keys.POSITIONS_KEY].reshape(-1, 3))
256256 # np.savetxt(os.path.join(out_dir, "atomic_numbers.dat"), structure[_keys.ATOMIC_NUMBERS_KEY], fmt='%d')
257257 # np.savetxt(os.path.join(out_dir, "pbc.dat"), structure[_keys.PBC_KEY])
258-
258+
259259 # write structure
260260 self .write_struct (structure , out_dir , fmt = fmt )
261261
@@ -266,6 +266,26 @@ def write_dat(self, idx, outroot, fmt='dat', eigenvalue=False, hamiltonian=False
266266 np .save (os .path .join (out_dir , "kpoints.npy" ), eigstatus [_keys .KPOINT_KEY ])
267267 np .save (os .path .join (out_dir , "eigenvalues.npy" ), eigstatus [_keys .ENERGY_EIGENVALUE_KEY ])
268268
269+ # write energy
270+ if energy :
271+ if hasattr (self , 'get_etot' ):
272+ energy_data = self .get_etot (idx )
273+ if energy_data is not None :
274+ np .savetxt (os .path .join (out_dir , "total_energy.dat" ), energy_data [_keys .TOTAL_ENERGY_KEY ])
275+
276+ # Write unconverged frame indices if present
277+ if _keys .UNCONVERGED_FRAME_INDICES_KEY in energy_data :
278+ unconverged_indices = energy_data [_keys .UNCONVERGED_FRAME_INDICES_KEY ]
279+ if len (unconverged_indices ) > 0 :
280+ with open (os .path .join (out_dir , "unconverged_frames.dat" ), 'w' ) as f :
281+ f .write ("# Frame indices that did not converge during MD/RELAX\n " )
282+ for idx_frame in unconverged_indices :
283+ f .write (f"{ idx_frame } \n " )
284+ else :
285+ log .warning (f"Failed to extract energy for structure { idx } " )
286+ else :
287+ log .warning (f"Parser does not implement get_etot method" )
288+
269289 # write blocks
270290 if any ([hamiltonian is not None , overlap is not None , density_matrix is not None ]) and any ([hamiltonian , overlap , density_matrix ]):
271291 with open (os .path .join (out_dir , "basis.dat" ), 'w' ) as f :
@@ -279,34 +299,55 @@ def write_dat(self, idx, outroot, fmt='dat', eigenvalue=False, hamiltonian=False
279299 for key_str , value in ham [i ].items ():
280300 default_group .create_dataset (key_str , data = value )
281301 del ham
282-
302+
283303 if overlap :
284304 with h5py .File (os .path .join (out_dir , "overlaps.h5" ), 'w' ) as fid :
285305 for i in range (len (ovp )):
286306 default_group = fid .create_group (str (i ))
287307 for key_str , value in ovp [i ].items ():
288308 default_group .create_dataset (key_str , data = value )
289309 del ovp
290-
310+
291311 if density_matrix :
292312 with h5py .File (os .path .join (out_dir , "density_matrices.h5" ), 'w' ) as fid :
293313 for i in range (len (dm )):
294314 default_group = fid .create_group (str (i ))
295315 for key_str , value in dm [i ].items ():
296316 default_group .create_dataset (key_str , data = value )
297-
317+
298318 del dm
299319
300320 return True
301321
302- def write_lmdb (self , idx , outroot , eigenvalue : bool = False , hamiltonian : bool = False , overlap : bool = False , density_matrix : bool = False ,band_index_min = 0 ):
322+ def write_lmdb (self , idx , outroot , eigenvalue : bool = False , hamiltonian : bool = False , overlap : bool = False , density_matrix : bool = False ,band_index_min = 0 , energy : bool = False ):
303323 os .makedirs (outroot , exist_ok = True )
304324 out_dir = os .path .join (outroot , "data.{}.lmdb" .format (os .getpid ()))
305325 structure = self .get_structure (idx )
306326 if any ([hamiltonian , overlap , density_matrix ]):
307327 ham , ovp , dm = self .get_blocks (idx , hamiltonian , overlap , density_matrix )
308328 if eigenvalue :
309329 eigstatus = self .get_eigenvalue (idx = idx , band_index_min = band_index_min )
330+ if energy :
331+ if hasattr (self , 'get_etot' ):
332+ energy_data = self .get_etot (idx )
333+ else :
334+ energy_data = None
335+ log .warning (f"Parser does not implement get_etot method" )
336+
337+ # Build frame index mapping for energy data
338+ # If there are unconverged frames, energy array will be shorter than n_frames
339+ energy_frame_mapping = None
340+ if energy and energy_data is not None :
341+ unconverged_indices = energy_data .get (_keys .UNCONVERGED_FRAME_INDICES_KEY , [])
342+ if len (unconverged_indices ) > 0 :
343+ # Build mapping: structure_frame_idx -> energy_array_idx
344+ energy_frame_mapping = {}
345+ energy_idx = 0
346+ n_frames_total = structure [_keys .POSITIONS_KEY ].shape [0 ]
347+ for frame_idx in range (n_frames_total ):
348+ if frame_idx not in unconverged_indices :
349+ energy_frame_mapping [frame_idx ] = energy_idx
350+ energy_idx += 1
310351
311352 n_frames = structure [_keys .POSITIONS_KEY ].shape [0 ]
312353 lmdb_env = lmdb .open (out_dir , map_size = 1048576000000 , lock = True )
@@ -321,6 +362,23 @@ def write_lmdb(self, idx, outroot, eigenvalue: bool=False, hamiltonian: bool=Fal
321362 data_dict [_keys .ENERGY_EIGENVALUE_KEY ] = eigstatus [_keys .ENERGY_EIGENVALUE_KEY ][nf ]
322363 data_dict [_keys .KPOINT_KEY ] = eigstatus [_keys .KPOINT_KEY ]
323364
365+ if energy and energy_data is not None :
366+ # For single structure (SCF/NSCF), energy_data has shape [1,]
367+ # For trajectories (MD/RELAX), energy_data has shape [nframes,] or less if unconverged
368+ if energy_data [_keys .TOTAL_ENERGY_KEY ].shape [0 ] == 1 :
369+ # Single structure case
370+ data_dict [_keys .TOTAL_ENERGY_KEY ] = energy_data [_keys .TOTAL_ENERGY_KEY ][0 ]
371+ else :
372+ # Trajectory case - use mapping if unconverged frames exist
373+ if energy_frame_mapping is not None :
374+ if nf in energy_frame_mapping :
375+ energy_idx = energy_frame_mapping [nf ]
376+ data_dict [_keys .TOTAL_ENERGY_KEY ] = energy_data [_keys .TOTAL_ENERGY_KEY ][energy_idx ]
377+ # else: skip energy for unconverged frames (don't add to data_dict)
378+ else :
379+ # No unconverged frames, direct indexing
380+ data_dict [_keys .TOTAL_ENERGY_KEY ] = energy_data [_keys .TOTAL_ENERGY_KEY ][nf ]
381+
324382 if hamiltonian :
325383 data_dict ["hamiltonian" ] = ham [nf ]
326384 if overlap :
0 commit comments