1+ #!/usr/bin/env python
2+ """Unified DeePKS command-line interface."""
3+
14import os
25import sys
3- import argparse
4- try :
5- import deepks
6- except ImportError as e :
7- sys .path .append (os .path .dirname (os .path .realpath (__file__ )) + "/../" )
8- from deepks .utils import load_yaml , deep_update
9-
10-
11- _MODEL_BACKEND = None
12- _PHYSICS_BACKEND = None
136
147
158def _get_model_backend ():
16- global _MODEL_BACKEND
17- if _MODEL_BACKEND is None :
18- from deepks .io .adapters import CorrNetModelBackend
19-
20- _MODEL_BACKEND = CorrNetModelBackend ()
21- return _MODEL_BACKEND
9+ """Get model backend instance."""
10+ from deepks .io .adapters import CorrNetModelBackend
11+ return CorrNetModelBackend ()
2212
2313
24- def _get_physics_backend ():
25- global _PHYSICS_BACKEND
26- if _PHYSICS_BACKEND is None :
27- from deepks .io .adapters import PySCFPhysicsBackend
14+ def _get_physics_backend (scf_soft = 'pyscf' ):
15+ """Get physics backend based on scf_soft parameter.
2816
29- _PHYSICS_BACKEND = PySCFPhysicsBackend ()
30- return _PHYSICS_BACKEND
17+ Args:
18+ scf_soft: SCF software name ('pyscf' or 'abacus')
3119
20+ Returns:
21+ PhysicsBackend instance
22+ """
23+ from deepks .core .physics import get_scf_backend
24+ return get_scf_backend (scf_soft )
3225
33- def main_cli (args = None ):
34- '''
35- Main function for DeepKS running. Call subfunctions to realize.
36- '''
37- # print("use modified deepks!")
38- parser = argparse .ArgumentParser (
39- prog = "deepks" ,
40- description = "A program to generate accurate energy functionals." )
41- parser .add_argument ("command" ,
42- help = "specify the sub-command to run, possible choices: "
43- "train, test, scf, stats, iterate" )
44- parser .add_argument ("args" , nargs = argparse .REMAINDER ,
45- help = "arguments to be passed to the sub-command" )
46-
47- args = parser .parse_args (args )
4826
49- # sepatate all sub_cli to make them useable independently
50- if args .command .upper () == "TRAIN" :
51- sub_cli = train_cli
52- elif args .command .upper () == "TEST" :
53- sub_cli = test_cli
54- elif args .command .upper () == "SCF" :
55- sub_cli = scf_cli
56- elif args .command .upper () == "STATS" :
57- sub_cli = stats_cli
58- elif args .command .upper ().startswith ("ITER" ):
59- sub_cli = iter_cli
60- else :
61- return ValueError (f"unsupported sub-command: { args .command } " )
62-
63- sub_cli (args .args )
27+ def main ():
28+ """Main entry point for DeePKS CLI."""
29+ import argparse
6430
65-
66- def train_cli (args = None ):
67- '''
68- Function for train process. Read parameters from command line or input yaml file.
69- '''
7031 parser = argparse .ArgumentParser (
71- prog = "deepks train" ,
72- description = "Train a model according to given input." ,
73- argument_default = argparse .SUPPRESS )
74- parser .add_argument ('input' , type = str , nargs = "?" ,
75- help = 'the input yaml file for args' )
76- parser .add_argument ('-r' , '--restart' ,
77- help = 'the restart file to load model from, would ignore model_args if given' )
78- parser .add_argument ('-d' , '--train-paths' , nargs = "*" ,
79- help = 'paths to the folders of training data' )
80- parser .add_argument ('-t' , '--test-paths' , nargs = "*" ,
81- help = 'paths to the folders of testing data' )
82- parser .add_argument ('-o' , '--ckpt-file' ,
83- help = 'file to save the model parameters, default: model.pth' )
84- parser .add_argument ("-P" , "--proj_basis" ,
85- help = "basis set used to project density matrix" )
86- parser .add_argument ('-S' , '--seed' , type = int ,
87- help = 'use specified seed in initialization and training' )
88- parser .add_argument ("-D" , "--device" ,
89- help = "device name used in training the model" )
90- args = parser .parse_args (args )
91-
92- if hasattr (args , "input" ):
93- argdict = load_yaml (args .input )
94- del args .input
95- argdict .update (vars (args ))
96- else :
97- argdict = vars (args )
98-
99- _get_model_backend ().train (** argdict )
100-
101-
102- def test_cli (args = None ):
103- '''
104- Function for test process. Read parameters from command line or input yaml file.
105- '''
106- parser = argparse .ArgumentParser (
107- prog = "deepks test" ,
108- description = "Test a model with given data (Not SCF)." ,
109- argument_default = argparse .SUPPRESS )
110- parser .add_argument ("input" , nargs = "?" ,
111- help = 'the input yaml file used for training' )
112- parser .add_argument ("-d" , "--data-paths" , type = str , nargs = '+' ,
113- help = "the paths to data folders containing .npy files for test" )
114- parser .add_argument ("-m" , "--model-file" , type = str , nargs = '+' ,
115- help = "the dumped model file to test" )
116- parser .add_argument ("-o" , "--output-prefix" , type = str ,
117- help = r"the prefix of output file, would wite into file %%prefix.%%sysidx.out" )
118- parser .add_argument ("-E" , "--e-name" , type = str ,
119- help = "the name of energy file to be read (no .npy extension)" )
120- parser .add_argument ("-D" , "--d-name" , type = str , nargs = "+" ,
121- help = "the name of descriptor file(s) to be read (no .npy extension)" )
122- parser .add_argument ("-G" , "--group" , action = 'store_true' ,
123- help = "group test results for all systems" )
124- args = parser .parse_args (args )
125-
126- if hasattr (args , "input" ):
127- rawdict = load_yaml (args .input )
128- del args .input
129- argdict = {}
130- if "ckpt_file" in rawdict ["train_args" ]:
131- argdict ["model_file" ] = rawdict ["train_args" ]["ckpt_file" ] # Check-point of the model
132- if "e_name" in rawdict ["data_args" ]:
133- argdict ["e_name" ] = rawdict ["data_args" ]["e_name" ]
134- if "d_name" in rawdict ["data_args" ]:
135- argdict ["d_name" ] = rawdict ["data_args" ]["d_name" ]
136- if "test_paths" in rawdict :
137- argdict ["data_paths" ] = rawdict ["test_paths" ]
138- argdict .update (vars (args ))
139- else :
140- argdict = vars (args )
141-
142- _get_model_backend ().evaluate (** argdict )
143-
144-
145- def scf_cli (args = None ):
146- '''
147- Function for calling scf procedure. Read parameters from command line or input yaml file.
148- '''
149- parser = argparse .ArgumentParser (
150- prog = "deepks scf" ,
151- description = "Calculate and save SCF results using given model." ,
152- argument_default = argparse .SUPPRESS )
153- parser .add_argument ("input" , nargs = "?" ,
154- help = 'the input yaml file for args' )
155- parser .add_argument ("-s" , "--systems" , nargs = "*" ,
156- help = "input molecule systems, can be xyz files or folders with npy data" )
157- parser .add_argument ("-m" , "--model-file" ,
158- help = "file of the trained model" )
159- parser .add_argument ("-d" , "--dump-dir" ,
160- help = "dir of dumped files" )
161- parser .add_argument ("-v" , "--verbose" , type = int , choices = range (0 ,6 ),
162- help = "output level of calculation information" )
163- parser .add_argument ("-F" , "--dump-fields" , nargs = "*" ,
164- help = "fields to be dumped into the folder" )
165- parser .add_argument ("-B" , "--basis" ,
166- help = "basis set used to solve the model" )
167- parser .add_argument ("-P" , "--proj_basis" ,
168- help = "basis set used to project dm, must match with model" )
169- parser .add_argument ("-D" , "--device" ,
170- help = "device name used in nn model inference" )
171- group0 = parser .add_mutually_exclusive_group ()
172- group0 .add_argument ("-G" , "--group" , action = 'store_true' , dest = "group" ,
173- help = "group results for all systems, only works for same number of atoms" )
174- group0 .add_argument ("-NG" , "--no-group" , action = 'store_false' , dest = "group" ,
175- help = "Do not group results for different systems (default behavior)" )
176- parser .add_argument ("-X" , "--scf-xc" ,
177- help = "base xc functional used in scf equation, default is HF" )
178- parser .add_argument ("--scf-conv-tol" , type = float ,
179- help = "converge threshold of scf iteration" )
180- parser .add_argument ("--scf-conv-tol-grad" , type = float ,
181- help = "gradient converge threshold of scf iteration" )
182- parser .add_argument ("--scf-max-cycle" , type = int ,
183- help = "max number of scf iteration cycles" )
184- parser .add_argument ("--scf-diis-space" , type = int ,
185- help = "subspace dimension used in diis mixing" )
186- parser .add_argument ("--scf-level-shift" , type = float ,
187- help = "level shift used in scf calculation" )
188-
189- args = parser .parse_args (args )
190-
191- scf_args = {}
192- # Combine parameters start with scf prefix
193- for k , v in vars (args ).copy ().items ():
194- if k .startswith ("scf_" ):
195- scf_args [k [4 :]] = v
196- delattr (args , k )
197-
198- if hasattr (args , "input" ):
199- argdict = load_yaml (args .input )
200- del args .input
201- argdict .update (vars (args ))
202- argdict ["scf_args" ].update (scf_args )
203- else :
204- argdict = vars (args )
205- argdict ["scf_args" ] = scf_args
206-
207- _get_physics_backend ().run_scf (** argdict )
208-
209-
210- def stats_cli (args = None ):
211- '''
212- Function for getting scf results. Read parameters from command line or input yaml file.
213- '''
214- parser = argparse .ArgumentParser (
215- prog = "deepks stats" ,
216- description = "Print the stats of SCF results." ,
217- argument_default = argparse .SUPPRESS )
218- parser .add_argument ("input" , nargs = "?" ,
219- help = 'the input yaml file used for SCF calculation' )
220- parser .add_argument ("-s" , "--systems" , nargs = "*" ,
221- help = 'system paths used as training set (i.e. calculate shift)' )
222- parser .add_argument ("-d" , "--dump-dir" ,
223- help = "directory used to save SCF results of training systems" )
224- parser .add_argument ("-ts" , "--test-sys" , nargs = "*" ,
225- help = 'system paths used as testing set (i.e. not calculate shift)' )
226- parser .add_argument ("-td" , "--test-dump" ,
227- help = "directory used to save SCF results of testing systems" )
228- parser .add_argument ("-G" , "--group" , action = 'store_true' ,
229- help = "if set, assume results are grouped" )
230- parser .add_argument ("-NC" , action = "store_false" , dest = "with_conv" ,
231- help = "do not print convergence results" )
232- parser .add_argument ("-NE" , action = "store_false" , dest = "with_e" ,
233- help = "do not print energy results" )
234- parser .add_argument ("-NF" , action = "store_false" , dest = "with_f" ,
235- help = "do not print force results" )
236- parser .add_argument ("--e-name" ,
237- help = "name of the energy file (no extension)" )
238- parser .add_argument ("--f-name" ,
239- help = "name of the force file (no extension)" )
240- args = parser .parse_args (args )
241-
242- if hasattr (args , "input" ):
243- rawdict = load_yaml (args .input )
244- del args .input
245- argdict = {fd : rawdict [fd ]
246- for fd in ("systems" , "dump_dir" , "group" )
247- if fd in rawdict }
248- argdict .update (vars (args ))
249- else :
250- argdict = vars (args )
251-
252- _get_physics_backend ().collect_stats (** argdict )
253-
254-
255- def iter_cli (args = None ):
256- '''
257- Function for doing iterations. Read parameters from command line or input yaml file.
258- '''
259- parser = argparse .ArgumentParser (
260- prog = "deepks iterate" ,
261- description = "Run the iteration procedure to train a SCF model." ,
262- argument_default = argparse .SUPPRESS )
263- parser .add_argument ("argfile" , nargs = "*" , default = [],
264- help = 'the input yaml file for args, '
265- 'if more than one, the latter has higher priority' )
266- parser .add_argument ("-s" , "--systems-train" , nargs = "*" ,
267- help = 'systems for training, '
268- 'can be xyz files or folders with npy data' )
269- parser .add_argument ("-t" , "--systems-test" , nargs = "*" ,
270- help = 'systems for training, '
271- 'can be xyz files or folders with npy data' )
272- parser .add_argument ("-n" , "--n-iter" , type = int ,
273- help = 'the number of iterations to run' )
274- parser .add_argument ("--workdir" ,
275- help = 'working directory, default is current directory' )
276- parser .add_argument ("--share-folder" ,
277- help = 'folder to store share files, default is "share"' )
278- parser .add_argument ("--cleanup" , action = "store_true" , dest = "cleanup" ,
279- help = 'if set, clean up files used for job dispatching' )
280- parser .add_argument ("--no-strict" , action = "store_false" , dest = "strict" ,
281- help = 'if set, allow other arguments to be passed to task' )
282- # allow cli specified argument files
283- sub_names = ["scf-input" , "scf-machine" , "train-input" , "train-machine" ,
284- "init-model" , "init-scf" , "init-train" , "scf-abacus" ]
285- for name in sub_names :
286- parser .add_argument (f"--{ name } " ,
287- help = 'if specified, subsitude the original arguments with given file' )
288-
289- args = parser .parse_args (args )
290- argdict = {}
291- for fl in args .argfile :
292- argdict = deep_update (argdict , load_yaml (fl ))
293- del args .argfile
294- argdict .update (vars (args ))
295-
296- from deepks .pipelines .iterate .iterate import main
297- main (** argdict )
32+ prog = "deepks" ,
33+ description = "DeePKS: Deep Kohn-Sham DFT with machine learning"
34+ )
35+ parser .add_argument (
36+ "config" ,
37+ nargs = "?" ,
38+ default = "input.yaml" ,
39+ help = "Configuration file (default: input.yaml)"
40+ )
41+ parser .add_argument (
42+ "-v" , "--version" ,
43+ action = "version" ,
44+ version = "DeePKS 1.0"
45+ )
46+
47+ args = parser .parse_args ()
48+
49+ # Check if config file exists
50+ if not os .path .exists (args .config ):
51+ print (f"Error: Configuration file '{ args .config } ' not found" , file = sys .stderr )
52+ sys .exit (1 )
53+
54+ # Load and process configuration
55+ from deepks .io .input import load_config , get_default_config
56+ from deepks .io .input .merger import merge_configs , apply_parameter_inheritance
57+ from deepks .io .input .dispatcher import dispatch_command
58+
59+ try :
60+ # Load configuration file
61+ config = load_config (args .config )
62+
63+ # Determine command from config
64+ if 'command' not in config :
65+ print ("Error: 'command' field is required in configuration file" , file = sys .stderr )
66+ print ("Valid commands: train, test, scf, stats, iterate" , file = sys .stderr )
67+ sys .exit (1 )
68+
69+ command = config ['command' ]
70+
71+ # Get defaults based on command and backend
72+ scf_soft = config .get ('scf_soft' , 'pyscf' )
73+ defaults = get_default_config (command , scf_soft )
74+
75+ # Merge defaults with config
76+ config = merge_configs (defaults , config )
77+
78+ # Apply parameter inheritance for iterate command
79+ if command == 'iterate' :
80+ config = apply_parameter_inheritance (config )
81+
82+ # Dispatch to appropriate handler
83+ dispatch_command (config )
84+
85+ except Exception as e :
86+ print (f"Error: { e } " , file = sys .stderr )
87+ import traceback
88+ traceback .print_exc ()
89+ sys .exit (1 )
29890
29991
30092if __name__ == "__main__" :
301- main_cli ()
93+ main ()
0 commit comments