Skip to content

Commit faafb4e

Browse files
Add eigenvalue solver option to structure calculations (#46)
* Feat: update elec_stru_cal.py as same as dptb, with numpy eigsolver choice * Feat: add eig_solver option to stru_options and update example JSON
1 parent 1c0c161 commit faafb4e

5 files changed

Lines changed: 91 additions & 32 deletions

File tree

dpnegf/runner/NEGF.py

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -178,11 +178,12 @@ def __init__(self,
178178
for lead_tag in ["lead_L", "lead_R"]:
179179
log.info(msg="-----Calculating Fermi level for {0}-----".format(lead_tag))
180180
_, e_fermi[lead_tag] = elec_cal.get_fermi_level(data=struct_leads[lead_tag],
181-
nel_atom = nel_atom_lead[lead_tag],
182-
meshgrid=self.stru_options[lead_tag]["kmesh_lead_Ef"],
183-
AtomicData_options=AtomicData_options,
184-
smearing_method=self.stru_options.get("e_fermi_smearing", "FD"),
185-
temp=100.0)
181+
nel_atom = nel_atom_lead[lead_tag],
182+
meshgrid=self.stru_options[lead_tag]["kmesh_lead_Ef"],
183+
AtomicData_options=AtomicData_options,
184+
smearing_method=self.stru_options.get("e_fermi_smearing", "FD"),
185+
temp=100.0,
186+
eig_solver=self.stru_options.get("eig_solver", "torch"),)
186187
else:
187188
e_fermi["lead_L"] = self.e_fermi
188189
e_fermi["lead_R"] = self.e_fermi

dpnegf/tests/test_get_fermi.py

Lines changed: 23 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,29 +1,37 @@
1-
import pytest
2-
from dpnegf.utils.elec_struc_cal import ElecStruCal
1+
from dptb.postprocess.elec_struc_cal import ElecStruCal
2+
3+
# from dptb.postprocess.bandstructure.band import Band
34
from dptb.nn.build import build_model
5+
import os
6+
from pathlib import Path
47

8+
rootdir = os.path.join(Path(os.path.abspath(__file__)).parent, "data")
59

6-
@pytest.fixture(scope='session', autouse=True)
7-
def root_directory(request):
8-
"""
9-
:return:
10-
"""
11-
return str(request.config.rootdir)
1210

13-
def test_get_fermi(root_directory):
14-
ckpt = f"{root_directory}/dpnegf/tests/data/test_get_fermi/nnsk.best.pth" # 'hopping': {'method': 'poly2exp', 'rs': 5.0, 'w': 0.6},
15-
stru_data = f"{root_directory}/dpnegf/tests/data/test_get_fermi/PRIMCELL.vasp"
11+
def test_get_fermi():
12+
ckpt = f"{rootdir}/test_get_fermi/nnsk.best.pth" # 'hopping': {'method': 'poly2exp', 'rs': 5.0, 'w': 0.6},
13+
stru_data = f"{rootdir}/test_get_fermi/PRIMCELL.vasp"
1614

1715
model = build_model(checkpoint=ckpt)
1816
nel_atom = {"Au":11}
1917

2018
elec_cal = ElecStruCal(model=model,device='cpu')
2119
_, efermi =elec_cal.get_fermi_level(data=stru_data,
2220
nel_atom = nel_atom,smearing_method='FD',
23-
meshgrid=[30,30,30])
24-
assert abs(efermi + 3.2257686853408813) < 1e-6
21+
meshgrid=[30,30,30],eig_solver='torch')
22+
assert abs(efermi + 3.2262574434280395) < 1e-3
23+
24+
_, efermi =elec_cal.get_fermi_level(data=stru_data,
25+
nel_atom = nel_atom,smearing_method='Gaussian',
26+
meshgrid=[30,30,30],eig_solver='torch')
27+
assert abs(efermi + 3.2262574434280395) < 1e-3
28+
29+
_, efermi =elec_cal.get_fermi_level(data=stru_data,
30+
nel_atom = nel_atom,smearing_method='FD',
31+
meshgrid=[30,30,30],eig_solver='numpy')
32+
assert abs(efermi + 3.2262574434280395) < 1e-3
2533

2634
_, efermi =elec_cal.get_fermi_level(data=stru_data,
2735
nel_atom = nel_atom,smearing_method='Gaussian',
28-
meshgrid=[30,30,30])
29-
assert abs(efermi + 3.2267462015151978) < 1e-6
36+
meshgrid=[30,30,30],eig_solver='numpy')
37+
assert abs(efermi + 3.2262574434280395) < 1e-3

dpnegf/utils/argcheck.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1076,6 +1076,7 @@ def stru_options():
10761076
doc_gamma_center=""
10771077
doc_time_reversal_symmetry=""
10781078
doc_e_fermi_smearing="The smearing method for Fermi level."
1079+
doc_eig_solver="The eigenvalue solver to use."
10791080
doc_nel_atom = "The number of electrons in each element."
10801081
return [
10811082
Argument("device", dict, optional=False, sub_fields=device(), doc=doc_device),
@@ -1086,6 +1087,7 @@ def stru_options():
10861087
Argument("gamma_center", list, optional=True, default=True, doc=doc_gamma_center),
10871088
Argument("time_reversal_symmetry", list, optional=True, default=True, doc=doc_time_reversal_symmetry),
10881089
Argument("e_fermi_smearing", str, optional=True, default="FD", doc=doc_e_fermi_smearing),
1090+
Argument("eig_solver", str, optional=True, default="torch", doc=doc_eig_solver),
10891091
Argument("nel_atom", [dict,None], optional=True, default=None, doc=doc_nel_atom)
10901092
]
10911093

dpnegf/utils/elec_struc_cal.py

Lines changed: 59 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,15 @@
1+
import os
2+
import h5py
13
import numpy as np
24
from ase.io import read
35
import ase
46
import numpy as np
57
from typing import Union
68
import torch
9+
from typing import Optional
710
import logging
811
log = logging.getLogger(__name__)
9-
from dptb.data import AtomicData, AtomicDataDict
12+
from dptb.data import AtomicData, AtomicDataDict, block_to_feature
1013
from dptb.nn.energy import Eigenvalues
1114
from dpnegf.utils.argcheck import get_cutoffs_from_model_options
1215
from copy import deepcopy
@@ -65,7 +68,12 @@ def __init__ (
6568
)
6669
r_max, er_max, oer_max = get_cutoffs_from_model_options(model.model_options)
6770
self.cutoffs = {'r_max': r_max, 'er_max': er_max, 'oer_max': oer_max}
68-
def get_data(self,data: Union[AtomicData, ase.Atoms, str],pbc:Union[bool,list]=None, device: Union[str, torch.device]=None, AtomicData_options:dict=None):
71+
def get_data(self,
72+
data: Union[AtomicData, ase.Atoms, str],
73+
pbc:Union[bool,list]=None,
74+
device: Union[str, torch.device]=None,
75+
AtomicData_options:dict=None,
76+
override_overlap:Optional[str]=None):
6977
'''The function `get_data` takes input data in the form of a string, ase.Atoms object, or AtomicData
7078
object, processes it accordingly, and returns the AtomicData class.
7179
@@ -80,6 +88,7 @@ def get_data(self,data: Union[AtomicData, ase.Atoms, str],pbc:Union[bool,list]=N
8088
device : Union[str, torch.device]
8189
The `device` parameter in the `get_data` function is used to specify the device on which the data
8290
should be processed. If no device is provided, it defaults to `self.device`.
91+
override_overlap : the path for overlap.h5 to use and override overlap matrix from model.
8392
8493
Returns
8594
-------
@@ -129,7 +138,30 @@ def get_data(self,data: Union[AtomicData, ase.Atoms, str],pbc:Union[bool,list]=N
129138
data = data
130139
else:
131140
raise ValueError('data should be either a string, ase.Atoms, or AtomicData')
132-
141+
142+
if isinstance(override_overlap, str):
143+
assert os.path.exists(override_overlap), "Overlap file not found."
144+
overlap_blocks = h5py.File(override_overlap, "r")
145+
if len(overlap_blocks) != 1:
146+
log.info('Overlap file contains more than one overlap matrix, only first will be used.')
147+
if self.overlap:
148+
log.warning('override_overlap is enabled while model contains overlap, override_overlap will be used.')
149+
if "0" in overlap_blocks:
150+
overlaps = overlap_blocks["0"]
151+
else:
152+
overlaps = overlap_blocks["1"]
153+
block_to_feature(data, self.model.idp, blocks=False, overlap_blocks=overlaps)
154+
if not self.overlap:
155+
self.eigv = Eigenvalues(
156+
idp=self.model.idp,
157+
device=self.device,
158+
s_edge_field=AtomicDataDict.EDGE_OVERLAP_KEY,
159+
s_node_field=AtomicDataDict.NODE_OVERLAP_KEY,
160+
s_out_field=AtomicDataDict.OVERLAP_KEY,
161+
dtype=self.model.dtype,
162+
)
163+
overlap_blocks.close()
164+
133165
if device is None:
134166
device = self.device
135167
data = AtomicData.to_AtomicDataDict(data.to(device))
@@ -138,7 +170,13 @@ def get_data(self,data: Union[AtomicData, ase.Atoms, str],pbc:Union[bool,list]=N
138170
return data
139171

140172

141-
def get_eigs(self, data: Union[AtomicData, ase.Atoms, str], klist: np.ndarray, pbc:Union[bool,list]=None, AtomicData_options:dict=None):
173+
def get_eigs(self,
174+
data: Union[AtomicData, ase.Atoms, str],
175+
klist: np.ndarray,
176+
pbc:Union[bool,list]=None,
177+
AtomicData_options:dict=None,
178+
override_overlap:Optional[str]=None,
179+
eig_solver:Optional[str]=None):
142180
'''This function calculates eigenvalues for Hk at specified k-points.
143181
144182
Parameters
@@ -151,28 +189,36 @@ def get_eigs(self, data: Union[AtomicData, ase.Atoms, str], klist: np.ndarray, p
151189
AtomicData_options : dict
152190
The `AtomicData_options` parameter is a dictionary that contains options for configuring the
153191
`AtomicData` object.
192+
override_overlap : the path for overlap.h5 to use and override overlap matrix from model.
154193
155194
Returns
156195
-------
157196
The function `get_eigs` returns the loaded data and the energy eigenvalues as a numpy array.
158197
159198
'''
160199

161-
data = self.get_data(data=data, pbc=pbc, device=self.device,AtomicData_options=AtomicData_options)
200+
data = self.get_data(data=data, pbc=pbc, device=self.device,AtomicData_options=AtomicData_options, override_overlap=override_overlap)
162201
# set the kpoint of the AtomicData
163202
data[AtomicDataDict.KPOINT_KEY] = \
164203
torch.nested.as_nested_tensor([torch.as_tensor(klist, dtype=self.model.dtype, device=self.device)])
204+
if isinstance(override_overlap, str):
205+
override_overlap_edge = data[AtomicDataDict.EDGE_OVERLAP_KEY]
206+
override_overlap_node = data[AtomicDataDict.NODE_OVERLAP_KEY]
165207
# get the eigenvalues
166208
data = self.model(data)
167-
if self.overlap == True:
209+
if isinstance(override_overlap, str):
210+
data[AtomicDataDict.EDGE_OVERLAP_KEY] = override_overlap_edge
211+
data[AtomicDataDict.NODE_OVERLAP_KEY] = override_overlap_node
212+
if self.overlap or isinstance(override_overlap, str):
168213
assert data.get(AtomicDataDict.EDGE_OVERLAP_KEY) is not None
169-
data = self.eigv(data)
214+
data = self.eigv(data, eig_solver=eig_solver)
170215

171216
return data, data[AtomicDataDict.ENERGY_EIGENVALUE_KEY][0].detach().cpu().numpy()
172217

173218
def get_fermi_level(self, data: Union[AtomicData, ase.Atoms, str], nel_atom: dict, \
174-
meshgrid: list = None, klist: np.ndarray=None, pbc:Union[bool,list]=None,AtomicData_options:dict=None,
175-
q_tol:float=1e-10,smearing_method:str='FD',temp:float=300):
219+
meshgrid: list = None, klist: np.ndarray=None, pbc:Union[bool,list]=None,
220+
AtomicData_options:dict=None, q_tol:float=1e-10, smearing_method:str='FD',
221+
temp:float=300,eig_solver:Optional[str]='torch'):
176222
'''This function calculates the Fermi level based on provided data with iteration method, electron counts per atom, and
177223
optional parameters like specific k-points and eigenvalues.
178224
@@ -233,13 +279,14 @@ def get_fermi_level(self, data: Union[AtomicData, ase.Atoms, str], nel_atom: dic
233279

234280
# eigenvalues would be used if provided, otherwise the eigenvalues would be calculated from the model on the specified k-points
235281
if not AtomicDataDict.ENERGY_EIGENVALUE_KEY in data:
236-
data, eigs = self.get_eigs(data=data, klist=klist, pbc=pbc, AtomicData_options=AtomicData_options)
282+
data, eigs = self.get_eigs(data=data, klist=klist, pbc=pbc,
283+
AtomicData_options=AtomicData_options,
284+
eig_solver=eig_solver)
237285
log.info('Getting eigenvalues from the model.')
238286
else:
239287
log.info('The eigenvalues are already in data. will use them.')
240288
eigs = data[AtomicDataDict.ENERGY_EIGENVALUE_KEY][0].detach().cpu().numpy()
241-
242-
289+
243290
if nel_atom is not None:
244291
atomtype_list = data[AtomicDataDict.ATOM_TYPE_KEY].flatten().tolist()
245292
atomtype_symbols = np.asarray(self.model.idp.type_names)[atomtype_list].tolist()

examples/atomic_chain_api/input_files/negf_chain_new.json

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@
1919
"time_reversal_symmetry": true,
2020
"nel_atom": {"C": 1.0},
2121
"kmesh":[1,1,1],
22+
"eig_solver":"numpy",
2223
"pbc":[false, false, false],
2324
"device":{
2425
"id":"4-8",

0 commit comments

Comments
 (0)