Skip to content

Commit 882ba94

Browse files
committed
refactor(adapters): wire cli through contract backends
1 parent f8e1365 commit 882ba94

4 files changed

Lines changed: 78 additions & 9 deletions

File tree

deepks/cli/main.py

Lines changed: 26 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,28 @@
88
from deepks.utils import load_yaml, deep_update
99

1010

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+
1133
def main_cli(args=None):
1234
'''
1335
Main function for DeepKS running. Call subfunctions to realize.
@@ -74,8 +96,7 @@ def train_cli(args=None):
7496
else:
7597
argdict = vars(args)
7698

77-
from deepks.pipelines.train.train import main
78-
main(**argdict)
99+
_get_model_backend().train(**argdict)
79100

80101

81102
def test_cli(args=None):
@@ -118,8 +139,7 @@ def test_cli(args=None):
118139
else:
119140
argdict = vars(args)
120141

121-
from deepks.pipelines.train.test import main
122-
main(**argdict)
142+
_get_model_backend().evaluate(**argdict)
123143

124144

125145
def scf_cli(args=None):
@@ -184,8 +204,7 @@ def scf_cli(args=None):
184204
argdict = vars(args)
185205
argdict["scf_args"] = scf_args
186206

187-
from deepks.pipelines.scf.run import main
188-
main(**argdict)
207+
_get_physics_backend().run_scf(**argdict)
189208

190209

191210
def stats_cli(args=None):
@@ -230,8 +249,7 @@ def stats_cli(args=None):
230249
else:
231250
argdict = vars(args)
232251

233-
from deepks.pipelines.scf.stats import print_stats
234-
print_stats(**argdict)
252+
_get_physics_backend().collect_stats(**argdict)
235253

236254

237255
def iter_cli(args=None):

deepks/io/adapters/__init__.py

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1,15 @@
1-
"""Scaffold package for refactor architecture."""
1+
"""Adapter implementations wiring contracts to concrete backends."""
2+
3+
__all__ = ["CorrNetModelBackend", "PySCFPhysicsBackend"]
4+
5+
6+
def __getattr__(name):
7+
if name == "CorrNetModelBackend":
8+
from .model_backend import CorrNetModelBackend
9+
10+
return CorrNetModelBackend
11+
if name == "PySCFPhysicsBackend":
12+
from .physics_backend import PySCFPhysicsBackend
13+
14+
return PySCFPhysicsBackend
15+
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
"""Model backend adapter implementations for CLI/pipeline integration."""
2+
3+
from deepks.core.contracts import ModelBackend
4+
from deepks.core.ml.eval.test import main as test_main
5+
from deepks.core.ml.train.train import main as train_main
6+
7+
8+
class CorrNetModelBackend(ModelBackend):
9+
"""Default ML backend wiring to current core ML entry points."""
10+
11+
def train(self, **kwargs):
12+
return train_main(**kwargs)
13+
14+
def evaluate(self, **kwargs):
15+
return test_main(**kwargs)
16+
17+
def predict(self, **kwargs):
18+
model = kwargs.pop("model")
19+
data = kwargs.pop("data")
20+
if kwargs:
21+
raise TypeError(f"unexpected predict kwargs: {sorted(kwargs.keys())}")
22+
return model(data)
Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,15 @@
1+
"""Physics backend adapter implementations for CLI/pipeline integration."""
2+
3+
from deepks.core.contracts import PhysicsBackend
4+
from deepks.core.physics.pyscf.run import main as scf_run_main
5+
from deepks.core.physics.pyscf.stats import print_stats
6+
7+
8+
class PySCFPhysicsBackend(PhysicsBackend):
9+
"""Default physics backend wiring to current PySCF-based entry points."""
10+
11+
def run_scf(self, **kwargs):
12+
return scf_run_main(**kwargs)
13+
14+
def collect_stats(self, **kwargs):
15+
return print_stats(**kwargs)

0 commit comments

Comments
 (0)