Skip to content

Commit 40f4b00

Browse files
Merge pull request #20 from AsymmetryChou/etot_add
Enable total energy extractor from ABACUS output
2 parents b788165 + e7f2c59 commit 40f4b00

38 files changed

Lines changed: 11124 additions & 35 deletions

.gitignore

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -157,3 +157,10 @@ cython_debug/
157157
# and can be added to the global gitignore or merged into this file. For a more nuclear
158158
# option (not recommended) you can uncomment the following to ignore the entire idea folder.
159159
#.idea/
160+
161+
162+
# test files
163+
test/data/siesta/siesta_out/*
164+
example/siesta_io/siesta_io.ipynb
165+
CLAUDE.md
166+
playground/*

dftio/__main__.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -153,6 +153,12 @@ def main_parser() -> argparse.ArgumentParser:
153153
default=0,
154154
help="The initial band index for eigenvalues to save.(0-band_index_min) bands will be ignored!"
155155
)
156+
parser_parse.add_argument(
157+
"-energy",
158+
"--energy",
159+
action="store_true",
160+
help="Whether to parse the total energy (Etot) from DFT output",
161+
)
156162

157163
parser_band = subparsers.add_parser(
158164
"band",

dftio/data/_keys.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -103,6 +103,7 @@
103103

104104
PER_ATOM_ENERGY_KEY: Final[str] = "atomic_energy"
105105
TOTAL_ENERGY_KEY: Final[str] = "total_energy"
106+
UNCONVERGED_FRAME_INDICES_KEY: Final[str] = "unconverged_frames"
106107
FORCE_KEY: Final[str] = "forces"
107108
PARTIAL_FORCE_KEY: Final[str] = "partial_forces"
108109
STRESS_KEY: Final[str] = "stress"

dftio/io/abacus/abacus_parser.py

Lines changed: 245 additions & 24 deletions
Large diffs are not rendered by default.

dftio/io/parse.py

Lines changed: 69 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -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:

dftio/io/vasp/vasp_parser.py

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -78,3 +78,34 @@ def read_EIGENVAL(file):
7878
def get_blocks(self, idx, hamiltonian: bool=False, overlap: bool=False, density_matrix: bool=False):
7979
raise NotImplementedError("VASP does not support block parsing yet.")
8080

81+
def get_total_energy(self, idx):
82+
path = self.raw_datas[idx]
83+
assert os.path.exists(os.path.join(path, "OUTCAR"))
84+
energy = self.read_total_energy(os.path.join(path, "OUTCAR"))
85+
return {_keys.TOTAL_ENERGY_KEY: np.array([energy], dtype=np.float64)}
86+
87+
# Alias for compatibility with base Parser class
88+
def get_etot(self, idx):
89+
"""Alias for get_total_energy to match base Parser convention."""
90+
log.warning("Only support for VASP static calculations. get_etot is an alias for get_total_energy.")
91+
return self.get_total_energy(idx)
92+
93+
@staticmethod
94+
def read_total_energy(file):
95+
"""
96+
Extract energy(sigma->0) from VASP OUTCAR file.
97+
This is the extrapolated energy to 0K.
98+
"""
99+
energy = []
100+
with open(file, 'r') as f:
101+
data = f.readlines()
102+
for line in data:
103+
if "energy(sigma->0)" in line:
104+
energy.append(float(re.findall(r'[\-\d\.E]+', line)[-1]))
105+
if len(energy) > 1:
106+
log.warning("Multiple energy(sigma->0) found in OUTCAR. Using the last one.")
107+
energy = energy[-1] if energy else None
108+
assert energy is not None, "Cannot find energy(sigma->0) in OUTCAR."
109+
110+
energy = np.array(energy, dtype=np.float64)
111+
return energy

pytest.ini

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
[pytest]
2+
# 让 pytest 优先导入当前目录(本地代码),而不是 site-packages
3+
pythonpath = .
4+
5+
# 避免 pytest 复制包目录导致导入旧版本
6+
addopts = --import-mode=importlib
7+

test/data/abacus_md/INPUT

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
INPUT_PARAMETERS
2+
#Parameters (1.General)
3+
calculation md
4+
nbands 20
5+
symmetry 0
6+
pseudo_dir ./
7+
orbital_dir ./
8+
9+
#Parameters (2.Iteration)
10+
ecutwfc 30
11+
scf_thr 1e-5
12+
scf_nmax 100
13+
14+
#Parameters (3.Basis)
15+
basis_type lcao
16+
ks_solver genelpa
17+
gamma_only 1
18+
19+
#Parameters (4.Smearing)
20+
smearing_method gaussian
21+
smearing_sigma 0.001
22+
23+
#Parameters (5.Mixing)
24+
mixing_type broyden
25+
mixing_beta 0.3
26+
chg_extrap second-order
27+
28+
#Parameters (6.MD)
29+
md_type nve
30+
md_nstep 10
31+
md_dt 1
32+
md_tfirst 300

test/data/abacus_md/KPT

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,4 @@
1+
K_POINTS
2+
0
3+
Gamma
4+
1 1 1 0 0 0

0 commit comments

Comments
 (0)