Skip to content

Commit a100a0e

Browse files
committed
update for sognet les_model
1 parent 04005de commit a100a0e

3 files changed

Lines changed: 35 additions & 75 deletions

File tree

deepmd/pt/model/model/__init__.py

Lines changed: 33 additions & 75 deletions
Original file line numberDiff line numberDiff line change
@@ -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-
363321
def 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__ = [

deepmd/pt/model/model/les_model.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -38,6 +38,7 @@
3838
@BaseModel.register("les_ener")
3939
class LESEnergyModel(DPModelCommon, LESEnergyModel_):
4040
model_type = "les_ener"
41+
fitting_net_type = "les_energy"
4142

4243
def __init__(
4344
self,

deepmd/pt/model/model/sog_model.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,7 @@
3939
@BaseModel.register("sog_ener")
4040
class SOGEnergyModel(DPModelCommon, SOGEnergyModel_):
4141
model_type = "sog_ener"
42+
fitting_net_type = "sog_energy"
4243

4344
def __init__(
4445
self,

0 commit comments

Comments
 (0)