|
81 | 81 | from deepmd.utils.data import ( |
82 | 82 | DataRequirementItem, |
83 | 83 | ) |
| 84 | +from deepmd.utils.finetune import ( |
| 85 | + warn_configuration_mismatch_during_finetune, |
| 86 | +) |
84 | 87 | from deepmd.utils.path import ( |
85 | 88 | DPH5Path, |
86 | 89 | ) |
@@ -520,9 +523,8 @@ def get_lr(lr_params: dict[str, Any]) -> BaseLR: |
520 | 523 | new_state_dict = {} |
521 | 524 | target_state_dict = self.wrapper.state_dict() |
522 | 525 | # pretrained_model |
523 | | - pretrained_model = get_model_for_wrapper( |
524 | | - state_dict["_extra_state"]["model_params"] |
525 | | - ) |
| 526 | + pretrained_model_params = state_dict["_extra_state"]["model_params"] |
| 527 | + pretrained_model = get_model_for_wrapper(pretrained_model_params) |
526 | 528 | pretrained_model_wrapper = ModelWrapper(pretrained_model) |
527 | 529 | pretrained_model_wrapper.set_state_dict(state_dict) |
528 | 530 | # update type related params |
@@ -557,6 +559,25 @@ def collect_single_finetune_params( |
557 | 559 | ) -> None: |
558 | 560 | _new_fitting = _finetune_rule_single.get_random_fitting() |
559 | 561 | _model_key_from = _finetune_rule_single.get_model_branch() |
| 562 | + _input_model_params = ( |
| 563 | + model_params["model_dict"][_model_key] |
| 564 | + if self.multi_task |
| 565 | + else model_params |
| 566 | + ) |
| 567 | + _pretrained_model_params = ( |
| 568 | + pretrained_model_params["model_dict"][_model_key_from] |
| 569 | + if "model_dict" in pretrained_model_params |
| 570 | + else pretrained_model_params |
| 571 | + ) |
| 572 | + if ( |
| 573 | + "descriptor" in _input_model_params |
| 574 | + and "descriptor" in _pretrained_model_params |
| 575 | + ): |
| 576 | + warn_configuration_mismatch_during_finetune( |
| 577 | + _input_model_params["descriptor"], |
| 578 | + _pretrained_model_params["descriptor"], |
| 579 | + _model_key_from, |
| 580 | + ) |
560 | 581 | target_keys = [ |
561 | 582 | i |
562 | 583 | for i in _random_state_dict.keys() |
|
0 commit comments