1+ import os
2+ import h5py
13import numpy as np
24from ase .io import read
35import ase
46import numpy as np
57from typing import Union
68import torch
9+ from typing import Optional
710import logging
811log = logging .getLogger (__name__ )
9- from dptb .data import AtomicData , AtomicDataDict
12+ from dptb .data import AtomicData , AtomicDataDict , block_to_feature
1013from dptb .nn .energy import Eigenvalues
1114from dpnegf .utils .argcheck import get_cutoffs_from_model_options
1215from 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 ()
0 commit comments