@@ -257,7 +257,10 @@ def _convert_preset_out_bias_to_array(
257257 return preset_out_bias
258258
259259
260- def get_standard_model (model_params : dict ) -> BaseModel :
260+ def get_standard_model (
261+ model_params : dict ,
262+ modelcls : type [BaseModel ] | None = None ,
263+ ) -> BaseModel :
261264 model_params_old = model_params
262265 model_params = copy .deepcopy (model_params )
263266 ntypes = len (model_params ["type_map" ])
@@ -272,63 +275,34 @@ def get_standard_model(model_params: dict) -> BaseModel:
272275 )
273276 data_stat_protect = model_params .get ("data_stat_protect" , 1e-2 )
274277
275- if fitting_net_type == "dipole" :
276- modelcls = DipoleModel
277- elif fitting_net_type == "polar" :
278- modelcls = PolarModel
279- elif fitting_net_type == "dos" :
280- modelcls = DOSModel
281- elif fitting_net_type in ["ener" , "direct_force_ener" ]:
282- modelcls = EnergyModel
283- elif fitting_net_type == "property" :
284- modelcls = PropertyModel
285- elif fitting_net_type == "sog_energy" :
286- modelcls = SOGEnergyModel
287- elif fitting_net_type == "les_energy" :
288- modelcls = LESEnergyModel
289- else :
290- raise RuntimeError (f"Unknown fitting type: { fitting_net_type } " )
291-
292- model = modelcls (
293- descriptor = descriptor ,
294- fitting = fitting ,
295- type_map = model_params ["type_map" ],
296- atom_exclude_types = atom_exclude_types ,
297- pair_exclude_types = pair_exclude_types ,
298- preset_out_bias = preset_out_bias ,
299- data_stat_protect = data_stat_protect ,
300- )
301- if model_params .get ("hessian_mode" ):
302- model .enable_hessian ()
303- model .model_def_script = json .dumps (model_params_old )
304- return model
305-
306-
307- def _get_lr_vmap_model (
308- model_params : dict ,
309- modelcls : type [BaseModel ],
310- expected_fitting_type : str ,
311- ) -> BaseModel :
312- model_params_old = model_params
313- model_params = copy .deepcopy (model_params )
314- ntypes = len (model_params ["type_map" ])
315- descriptor , fitting , fitting_net_type = _get_standard_model_components (
316- model_params , ntypes
317- )
318- if fitting_net_type != expected_fitting_type :
278+ if modelcls is None :
279+ if fitting_net_type == "dipole" :
280+ modelcls = DipoleModel
281+ elif fitting_net_type == "polar" :
282+ modelcls = PolarModel
283+ elif fitting_net_type == "dos" :
284+ modelcls = DOSModel
285+ elif fitting_net_type in ["ener" , "direct_force_ener" ]:
286+ modelcls = EnergyModel
287+ elif fitting_net_type == "property" :
288+ modelcls = PropertyModel
289+ else :
290+ # Auto-discover from BaseModel plugins
291+ for plugin_cls in BaseModel .get_plugins ().values ():
292+ if getattr (plugin_cls , "fitting_net_type" , None ) == fitting_net_type :
293+ modelcls = plugin_cls
294+ break
295+ if modelcls is None :
296+ raise RuntimeError (f"Unknown fitting type: { fitting_net_type } " )
297+
298+ # Validate fitting type when modelcls is explicitly provided (e.g. vmap variants)
299+ expected_fitting = getattr (modelcls , "fitting_net_type" , None )
300+ if expected_fitting is not None and fitting_net_type != expected_fitting :
319301 raise RuntimeError (
320- f"{ modelcls .__name__ } requires fitting_net.type='{ expected_fitting_type } ', "
302+ f"{ modelcls .__name__ } requires fitting_net.type='{ expected_fitting } ', "
321303 f"got '{ fitting_net_type } '."
322304 )
323305
324- atom_exclude_types = model_params .get ("atom_exclude_types" , [])
325- pair_exclude_types = model_params .get ("pair_exclude_types" , [])
326- preset_out_bias = model_params .get ("preset_out_bias" )
327- preset_out_bias = _convert_preset_out_bias_to_array (
328- preset_out_bias , model_params ["type_map" ]
329- )
330- data_stat_protect = model_params .get ("data_stat_protect" , 1e-2 )
331-
332306 model = modelcls (
333307 descriptor = descriptor ,
334308 fitting = fitting ,
@@ -344,22 +318,6 @@ def _get_lr_vmap_model(
344318 return model
345319
346320
347- def get_sog_vmap_model (model_params : dict ) -> BaseModel :
348- return _get_lr_vmap_model (
349- model_params ,
350- modelcls = SOGVmapModel ,
351- expected_fitting_type = "sog_energy" ,
352- )
353-
354-
355- def get_les_vmap_model (model_params : dict ) -> BaseModel :
356- return _get_lr_vmap_model (
357- model_params ,
358- modelcls = LESVmapModel ,
359- expected_fitting_type = "les_energy" ,
360- )
361-
362-
363321def get_model (model_params : dict ) -> Any :
364322 model_type = model_params .get ("type" , "standard" )
365323 if model_type == "standard" :
@@ -369,14 +327,14 @@ def get_model(model_params: dict) -> Any:
369327 return get_zbl_model (model_params )
370328 else :
371329 return get_standard_model (model_params )
372- elif model_type == "sog_vmap" :
373- return get_sog_vmap_model (model_params )
374- elif model_type == "les_vmap" :
375- return get_les_vmap_model (model_params )
376330 elif model_type == "linear_ener" :
377331 return get_linear_model (model_params )
378332 else :
379- return BaseModel .get_class_by_type (model_type ).get_model (model_params )
333+ plugin_cls = BaseModel .get_class_by_type (model_type )
334+ if hasattr (plugin_cls , "get_model" ) and callable (plugin_cls .get_model ):
335+ return plugin_cls .get_model (model_params )
336+ # Fallback: reuse standard model construction logic
337+ return get_standard_model (model_params , modelcls = plugin_cls )
380338
381339
382340__all__ = [
0 commit comments