1- import os ,time ,sys
1+ import os
2+ import sys
23import numpy as np
34import torch
4- from deepks .model .utils import make_integrator , cal_nb_overlap
5-
6- def concat_batch (tdicts , dim = 0 ):
7- keys = tdicts [0 ].keys ()
8- assert all (d .keys () == keys for d in tdicts )
9- return {
10- k : torch .cat ([d [k ] for d in tdicts ], dim )
11- for k in keys
12- }
13-
14- def split_batch (tdict , size , dim = 0 , global_keys = None ):
15- if global_keys is None :
16- global_keys = {"data_shape" }
17- dsplit = {}
18- for k ,v in tdict .items ():
19- if k in global_keys :
20- dsplit [k ] = v
21- elif isinstance (v , torch .Tensor ):
22- dsplit [k ] = torch .split (v , size , dim )
23- elif isinstance (v , np .ndarray ):
24- # support only for dim=0
25- assert dim == 0 , "numpy.ndarray supports only for dim=0 split"
26- dsplit [k ] = np .array_split (v , range (size , v .shape [0 ], size ), axis = 0 )
27- elif isinstance (v , list ):
28- # support only for dim=0
29- assert dim == 0 , "list supports only for dim=0 split"
30- dsplit [k ] = [v [i :i + size ] for i in range (0 , len (v ), size )]
31- else :
32- raise TypeError (f"Unsupported type for split_batch: { type (v )} " )
33- # dsplit = {k: torch.split(v, size, dim) for k,v in tdict.items()}
34- nsecs = [len (v ) for k , v in dsplit .items () if k not in global_keys ]
35- assert all (ns == nsecs [0 ] for ns in nsecs )
36- return [
37- {k : (v [i ] if k not in global_keys else v ) for k , v in dsplit .items ()}
38- for i in range (nsecs [0 ])
39- ]
405
41- def generalized_eigh (h ,L_inv ):
42- symm_h = L_inv @ h @ L_inv .mT
43- e ,v = torch .linalg .eigh (symm_h )
44- phi = L_inv .mT @ v
45- return e ,phi
6+ from deepks .io .schemas .reader_fields import (
7+ DEFAULT_READER_FIELD_NAMES ,
8+ ReaderFieldNames ,
9+ resolve_reader_paths ,
10+ )
11+ from deepks .io .transforms .batch import concat_batch , split_batch
12+ from deepks .io .transforms .linalg import generalized_eigh
13+ from deepks .model .utils import make_integrator , cal_nb_overlap
4614
4715class Reader (object ):
4816 def __init__ (self , data_path , batch_size ,
49- e_name = "l_e_delta" , d_name = "dm_eig" ,
50- f_name = "l_f_delta" , gvx_name = "grad_vx" ,
51- s_name = "l_s_delta" , gvepsl_name = "grad_vepsl" ,
52- o_name = "l_o_delta" , op_name = "orbital_precalc" ,
53- h_name = "l_h_delta" , vdp_name = "v_delta_precalc" ,
54- vdrp_name = "vdr_precalc" , phialpha_name = "phialpha" ,
55- gevdm_name = "grad_evdm" , hr_name = "l_hr_delta" ,
56- h_base_name = "h_base" , h_ref_name = "hamiltonian" ,
57- read_overlap = False , overlap_name = "overlap" ,
58- eg_name = "eg_base" , gveg_name = "grad_veg" ,
59- gldv_name = "grad_ldv" , conv_name = "conv" ,
60- atom_name = "atom" , box_name = "box" ,
17+ e_name = DEFAULT_READER_FIELD_NAMES .e_name ,
18+ d_name = DEFAULT_READER_FIELD_NAMES .d_name ,
19+ f_name = DEFAULT_READER_FIELD_NAMES .f_name ,
20+ gvx_name = DEFAULT_READER_FIELD_NAMES .gvx_name ,
21+ s_name = DEFAULT_READER_FIELD_NAMES .s_name ,
22+ gvepsl_name = DEFAULT_READER_FIELD_NAMES .gvepsl_name ,
23+ o_name = DEFAULT_READER_FIELD_NAMES .o_name ,
24+ op_name = DEFAULT_READER_FIELD_NAMES .op_name ,
25+ h_name = DEFAULT_READER_FIELD_NAMES .h_name ,
26+ vdp_name = DEFAULT_READER_FIELD_NAMES .vdp_name ,
27+ vdrp_name = DEFAULT_READER_FIELD_NAMES .vdrp_name ,
28+ phialpha_name = DEFAULT_READER_FIELD_NAMES .phialpha_name ,
29+ gevdm_name = DEFAULT_READER_FIELD_NAMES .gevdm_name ,
30+ hr_name = DEFAULT_READER_FIELD_NAMES .hr_name ,
31+ h_base_name = DEFAULT_READER_FIELD_NAMES .h_base_name ,
32+ h_ref_name = DEFAULT_READER_FIELD_NAMES .h_ref_name ,
33+ read_overlap = False ,
34+ overlap_name = DEFAULT_READER_FIELD_NAMES .overlap_name ,
35+ eg_name = DEFAULT_READER_FIELD_NAMES .eg_name ,
36+ gveg_name = DEFAULT_READER_FIELD_NAMES .gveg_name ,
37+ gldv_name = DEFAULT_READER_FIELD_NAMES .gldv_name ,
38+ conv_name = DEFAULT_READER_FIELD_NAMES .conv_name ,
39+ atom_name = DEFAULT_READER_FIELD_NAMES .atom_name ,
40+ box_name = DEFAULT_READER_FIELD_NAMES .box_name ,
6141 orb_list = None , alpha_list = None , ** kwargs ):
6242 self .data_path = data_path
6343 self .batch_size = batch_size
64- self .e_path = self .check_exist (e_name + ".npy" )
65- self .f_path = self .check_exist (f_name + ".npy" )
66- self .s_path = self .check_exist (s_name + ".npy" )
67- self .o_path = self .check_exist (o_name + ".npy" )
68- self .h_path = self .check_exist (h_name + ".npy" )
69- self .hr_path = self .check_exist (hr_name + ".npy" )
70- self .h_base_path = self .check_exist (h_base_name + ".npy" )
71- self .h_ref_path = self .check_exist (h_ref_name + ".npy" )
72- self .overlap_path = self .check_exist (overlap_name + ".npy" )
73- self .phialpha_path = self .check_exist (phialpha_name + ".npy" )
74- self .gevdm_path = self .check_exist (gevdm_name + ".npy" )
75- self .d_path = self .check_exist (d_name + ".npy" )
76- self .gvx_path = self .check_exist (gvx_name + ".npy" )
77- self .gvepsl_path = self .check_exist (gvepsl_name + ".npy" )
78- self .op_path = self .check_exist (op_name + ".npy" )
79- self .vdp_path = self .check_exist (vdp_name + ".npy" )
80- self .vdrp_path = self .check_exist (vdrp_name + ".npy" )
81- self .eg_path = self .check_exist (eg_name + ".npy" )
82- self .gveg_path = self .check_exist (gveg_name + ".npy" )
83- self .gldv_path = self .check_exist (gldv_name + ".npy" )
84- self .c_path = self .check_exist (conv_name + ".npy" )
85- self .a_path = self .check_exist (atom_name + ".npy" )
86- self .b_path = self .check_exist (box_name + ".npy" )
44+ field_names = ReaderFieldNames (
45+ e_name = e_name ,
46+ d_name = d_name ,
47+ f_name = f_name ,
48+ gvx_name = gvx_name ,
49+ s_name = s_name ,
50+ gvepsl_name = gvepsl_name ,
51+ o_name = o_name ,
52+ op_name = op_name ,
53+ h_name = h_name ,
54+ vdp_name = vdp_name ,
55+ vdrp_name = vdrp_name ,
56+ phialpha_name = phialpha_name ,
57+ gevdm_name = gevdm_name ,
58+ hr_name = hr_name ,
59+ h_base_name = h_base_name ,
60+ h_ref_name = h_ref_name ,
61+ overlap_name = overlap_name ,
62+ eg_name = eg_name ,
63+ gveg_name = gveg_name ,
64+ gldv_name = gldv_name ,
65+ conv_name = conv_name ,
66+ atom_name = atom_name ,
67+ box_name = box_name ,
68+ )
69+ for path_name , path in resolve_reader_paths (self .data_path , field_names ).items ():
70+ setattr (self , path_name , path )
71+
72+ self .system_raw_path = os .path .join (self .data_path , "system.raw" )
8773 self .read_overlap = read_overlap
8874 self .orb_list = ["../../" + orb for orb in orb_list ] if orb_list is not None else None
8975 self .alpha_list = ["../../" + alpha for alpha in alpha_list ] if alpha_list is not None else None
@@ -93,16 +79,9 @@ def __init__(self, data_path, batch_size,
9379 # initialize sample index queue
9480 self .idx_queue = []
9581
96- def check_exist (self , fname ):
97- if fname is None :
98- return None
99- fpath = os .path .join (self .data_path , fname )
100- if os .path .exists (fpath ):
101- return fpath
102-
10382 def load_meta (self ):
10483 try :
105- sys_meta = np .loadtxt (self .check_exist ( 'system.raw' ) , converters = float ).astype (int ).reshape ([- 1 ])
84+ sys_meta = np .loadtxt (self .system_raw_path , converters = float ).astype (int ).reshape ([- 1 ])
10685 self .natm = sys_meta [0 ]
10786 self .nproj = sys_meta [- 1 ]
10887 except :
@@ -526,8 +505,11 @@ def revert_elem_const(self):
526505
527506class SimpleReader (object ):
528507 def __init__ (self , data_path , batch_size ,
529- e_name = "l_e_delta" , d_name = "dm_eig" ,
530- conv_filter = True , conv_name = "conv" , ** kwargs ):
508+ e_name = DEFAULT_READER_FIELD_NAMES .e_name ,
509+ d_name = DEFAULT_READER_FIELD_NAMES .d_name ,
510+ conv_filter = True ,
511+ conv_name = DEFAULT_READER_FIELD_NAMES .conv_name ,
512+ ** kwargs ):
531513 # copy from config
532514 self .data_path = data_path
533515 self .batch_size = batch_size
0 commit comments