Skip to content

Commit 65d37b8

Browse files
committed
refactor: unified CLI with config-driven architecture
Major Changes: - Replace multiple CLI entry points with single unified entry - Implement config-driven command dispatch system - Add unified I/O input module with config loading, merging, validation - Abstract SCF backend interface supporting PySCF and ABACUS - Add parameter inheritance for iterate command - Remove deprecated use_abacus parameter, use scf_soft instead New Modules: - deepks/io/input/: Unified configuration processing - defaults.py: Default configuration values - loader.py: Configuration file loading - merger.py: Deep config merging and parameter inheritance - validator.py: Configuration validation - dispatcher.py: Command dispatching - deepks/core/physics/factory.py: Backend factory - deepks/core/physics/abacus/run.py: ABACUS backend implementation CLI Changes: - Single entry point: python deepks/cli/main.py [config.yaml] - Default config file: input.yaml - Config must specify 'command' field (train/test/scf/stats/iterate) - Backend selection via 'scf_soft' parameter (pyscf/abacus) Breaking Changes: - Removed old CLI functions (train_cli, test_cli, scf_cli, etc.) - Removed use_abacus parameter (replaced by scf_soft) - Configuration files must include 'command' field Tests: - All 101 tests passing - Updated smoke tests for unified CLI - Added 26 new tests for I/O module and unified CLI Documentation: - Added docs/input-parameter.md with complete parameter reference
1 parent 5f05a76 commit 65d37b8

24 files changed

Lines changed: 3362 additions & 500 deletions

deepks/cli/main.py

Lines changed: 77 additions & 285 deletions
Original file line numberDiff line numberDiff line change
@@ -1,301 +1,93 @@
1+
#!/usr/bin/env python
2+
"""Unified DeePKS command-line interface."""
3+
14
import os
25
import 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

158
def _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

30092
if __name__ == "__main__":
301-
main_cli()
93+
main()

0 commit comments

Comments
 (0)