|
8 | 8 | from deepks.utils import load_yaml, deep_update |
9 | 9 |
|
10 | 10 |
|
| 11 | +_MODEL_BACKEND = None |
| 12 | +_PHYSICS_BACKEND = None |
| 13 | + |
| 14 | + |
| 15 | +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 |
| 22 | + |
| 23 | + |
| 24 | +def _get_physics_backend(): |
| 25 | + global _PHYSICS_BACKEND |
| 26 | + if _PHYSICS_BACKEND is None: |
| 27 | + from deepks.io.adapters import PySCFPhysicsBackend |
| 28 | + |
| 29 | + _PHYSICS_BACKEND = PySCFPhysicsBackend() |
| 30 | + return _PHYSICS_BACKEND |
| 31 | + |
| 32 | + |
11 | 33 | def main_cli(args=None): |
12 | 34 | ''' |
13 | 35 | Main function for DeepKS running. Call subfunctions to realize. |
@@ -74,8 +96,7 @@ def train_cli(args=None): |
74 | 96 | else: |
75 | 97 | argdict = vars(args) |
76 | 98 |
|
77 | | - from deepks.pipelines.train.train import main |
78 | | - main(**argdict) |
| 99 | + _get_model_backend().train(**argdict) |
79 | 100 |
|
80 | 101 |
|
81 | 102 | def test_cli(args=None): |
@@ -118,8 +139,7 @@ def test_cli(args=None): |
118 | 139 | else: |
119 | 140 | argdict = vars(args) |
120 | 141 |
|
121 | | - from deepks.pipelines.train.test import main |
122 | | - main(**argdict) |
| 142 | + _get_model_backend().evaluate(**argdict) |
123 | 143 |
|
124 | 144 |
|
125 | 145 | def scf_cli(args=None): |
@@ -184,8 +204,7 @@ def scf_cli(args=None): |
184 | 204 | argdict = vars(args) |
185 | 205 | argdict["scf_args"] = scf_args |
186 | 206 |
|
187 | | - from deepks.pipelines.scf.run import main |
188 | | - main(**argdict) |
| 207 | + _get_physics_backend().run_scf(**argdict) |
189 | 208 |
|
190 | 209 |
|
191 | 210 | def stats_cli(args=None): |
@@ -230,8 +249,7 @@ def stats_cli(args=None): |
230 | 249 | else: |
231 | 250 | argdict = vars(args) |
232 | 251 |
|
233 | | - from deepks.pipelines.scf.stats import print_stats |
234 | | - print_stats(**argdict) |
| 252 | + _get_physics_backend().collect_stats(**argdict) |
235 | 253 |
|
236 | 254 |
|
237 | 255 | def iter_cli(args=None): |
|
0 commit comments