11# SPDX-License-Identifier: LGPL-3.0-or-later
22import copy
3- from typing import (
4- Any ,
5- )
63
74from deepmd .dpmodel .atomic_model .dp_atomic_model import (
85 DPAtomicModel ,
1613from deepmd .dpmodel .fitting .base_fitting import (
1714 BaseFitting ,
1815)
19- from deepmd .dpmodel .fitting . ener_fitting import (
20- EnergyFittingNet ,
21- )
16+ from deepmd .dpmodel .model import dipole_model as _dipole_model # noqa: F401
17+ from deepmd . dpmodel . model import dos_model as _dos_model # noqa: F401
18+ from deepmd . dpmodel . model import polar_model as _polar_model # noqa: F401
2219from deepmd .dpmodel .model .base_model import (
2320 BaseModel ,
2421)
25- from deepmd .dpmodel .model .dipole_model import (
26- DipoleModel ,
27- )
28- from deepmd .dpmodel .model .dos_model import (
29- DOSModel ,
30- )
3122from deepmd .dpmodel .model .dp_zbl_model import (
3223 DPZBLModel ,
3324)
3425from deepmd .dpmodel .model .ener_model import (
3526 EnergyModel ,
3627)
37- from deepmd .dpmodel .model .polar_model import (
38- PolarModel ,
28+ from deepmd .dpmodel .model .model_factory import get_model as get_model_from_factory
29+ from deepmd .dpmodel .model .model_factory import (
30+ get_standard_model as get_standard_model_from_factory ,
3931)
40- from deepmd .dpmodel .model .property_model import (
41- PropertyModel ,
32+ from deepmd .dpmodel .model .model_factory import (
33+ get_zbl_model as get_zbl_model_from_factory ,
4234)
4335from deepmd .dpmodel .model .spin_model import (
4436 SpinModel ,
4840)
4941
5042
51- def _get_standard_model_components (
52- data : dict [str , Any ], ntypes : int
53- ) -> tuple [BaseDescriptor , BaseFitting , str ]:
54- # descriptor
55- data ["descriptor" ]["ntypes" ] = ntypes
56- data ["descriptor" ]["type_map" ] = copy .deepcopy (data ["type_map" ])
57- descriptor = BaseDescriptor (** data ["descriptor" ])
58- # fitting
59- fitting_net = data .get ("fitting_net" , {})
60- fitting_net ["type" ] = fitting_net .get ("type" , "ener" )
61- fitting_net ["ntypes" ] = descriptor .get_ntypes ()
62- fitting_net ["type_map" ] = copy .deepcopy (data ["type_map" ])
63- fitting_net ["mixed_types" ] = descriptor .mixed_types ()
64- if fitting_net ["type" ] in ["dipole" , "polar" ]:
65- fitting_net ["embedding_width" ] = descriptor .get_dim_emb ()
66- fitting_net ["dim_descrpt" ] = descriptor .get_dim_out ()
67- grad_force = "direct" not in fitting_net ["type" ]
68- if not grad_force :
69- fitting_net ["out_dim" ] = descriptor .get_dim_emb ()
70- if "ener" in fitting_net ["type" ]:
71- fitting_net ["return_energy" ] = True
72- fitting = BaseFitting (** fitting_net )
73- return descriptor , fitting , fitting_net ["type" ]
74-
75-
7643def get_standard_model (data : dict ) -> EnergyModel :
7744 """Get a EnergyModel from a dictionary.
7845
@@ -81,78 +48,24 @@ def get_standard_model(data: dict) -> EnergyModel:
8148 data : dict
8249 The data to construct the model.
8350 """
84- if "type_embedding" in data :
85- raise ValueError (
86- "In the DP backend, type_embedding is not at the model level, but within the descriptor. See type embedding documentation for details."
87- )
88- data = copy .deepcopy (data )
89- ntypes = len (data ["type_map" ])
90- descriptor , fitting , fitting_net_type = _get_standard_model_components (data , ntypes )
91- atom_exclude_types = data .get ("atom_exclude_types" , [])
92- pair_exclude_types = data .get ("pair_exclude_types" , [])
93-
94- if fitting_net_type == "dipole" :
95- modelcls = DipoleModel
96- elif fitting_net_type == "polar" :
97- modelcls = PolarModel
98- elif fitting_net_type == "dos" :
99- modelcls = DOSModel
100- elif fitting_net_type in ["ener" , "direct_force_ener" ]:
101- modelcls = EnergyModel
102- elif fitting_net_type == "property" :
103- modelcls = PropertyModel
104- else :
105- raise RuntimeError (f"Unknown fitting type: { fitting_net_type } " )
106-
107- model = modelcls (
108- descriptor = descriptor ,
109- fitting = fitting ,
110- type_map = data ["type_map" ],
111- atom_exclude_types = atom_exclude_types ,
112- pair_exclude_types = pair_exclude_types ,
51+ return get_standard_model_from_factory (
52+ data ,
53+ descriptor_base = BaseDescriptor ,
54+ fitting_base = BaseFitting ,
55+ model_base = BaseModel ,
56+ backend_name = "DP" ,
11357 )
114- return model
11558
11659
11760def get_zbl_model (data : dict ) -> DPZBLModel :
118- data = copy .deepcopy (data )
119- data ["descriptor" ]["ntypes" ] = len (data ["type_map" ])
120- data ["descriptor" ]["type_map" ] = data ["type_map" ]
121- descriptor = BaseDescriptor (** data ["descriptor" ])
122- fitting_type = data ["fitting_net" ].pop ("type" )
123- data ["fitting_net" ]["type_map" ] = data ["type_map" ]
124- if fitting_type == "ener" :
125- fitting = EnergyFittingNet (
126- ntypes = descriptor .get_ntypes (),
127- dim_descrpt = descriptor .get_dim_out (),
128- mixed_types = descriptor .mixed_types (),
129- ** data ["fitting_net" ],
130- )
131- else :
132- raise ValueError (f"Unknown fitting type { fitting_type } " )
133-
134- dp_model = DPAtomicModel (descriptor , fitting , type_map = data ["type_map" ])
135- # pairtab
136- filepath = data ["use_srtab" ]
137- pt_model = PairTabAtomicModel (
138- filepath ,
139- descriptor .get_rcut (),
140- descriptor .get_sel (),
141- type_map = data ["type_map" ],
142- )
143-
144- rmin = data ["sw_rmin" ]
145- rmax = data ["sw_rmax" ]
146- atom_exclude_types = data .get ("atom_exclude_types" , [])
147- pair_exclude_types = data .get ("pair_exclude_types" , [])
148- return DPZBLModel (
149- dp_model ,
150- pt_model ,
151- rmin ,
152- rmax ,
153- type_map = data ["type_map" ],
154- atom_exclude_types = atom_exclude_types ,
155- pair_exclude_types = pair_exclude_types ,
61+ return get_zbl_model_from_factory (
62+ data ,
63+ descriptor_base = BaseDescriptor ,
64+ fitting_base = BaseFitting ,
65+ atomic_model = DPAtomicModel ,
66+ pairtab_model = PairTabAtomicModel ,
67+ zbl_model = DPZBLModel ,
68+ backend_name = "DP" ,
15669 )
15770
15871
@@ -198,13 +111,10 @@ def get_model(data: dict) -> BaseModel:
198111 data : dict
199112 The data to construct the model.
200113 """
201- model_type = data .get ("type" , "standard" )
202- if model_type == "standard" :
203- if "spin" in data :
204- return get_spin_model (data )
205- elif "use_srtab" in data :
206- return get_zbl_model (data )
207- else :
208- return get_standard_model (data )
209- else :
210- return BaseModel .get_class_by_type (model_type ).get_model (data )
114+ return get_model_from_factory (
115+ data ,
116+ base_model = BaseModel ,
117+ standard_model_factory = get_standard_model ,
118+ spin_model_factory = get_spin_model ,
119+ zbl_model_factory = get_zbl_model ,
120+ )
0 commit comments