diff --git a/deepmd/pt/train/training.py b/deepmd/pt/train/training.py index 8dd68c8cdb..85c4fa76db 100644 --- a/deepmd/pt/train/training.py +++ b/deepmd/pt/train/training.py @@ -1616,26 +1616,20 @@ def fake_model() -> dict: def log_loss_train( _loss: Any, _more_loss: Any, _task_key: str = "Default" ) -> dict: - results = {} if not self.multi_task: - # Use accumulated average loss for single task - for item in self.train_loss_accu: - results[item] = ( - self.train_loss_accu[item] - / self.step_count_in_interval - ) - else: - # Use accumulated average loss for multi-task - if ( - _task_key in self.train_loss_accu - and _task_key in self.step_count_per_task - ): - for item in self.train_loss_accu[_task_key]: - results[item] = ( - self.train_loss_accu[_task_key][item] - / self.step_count_per_task[_task_key] - ) - return results + return { + item: value / self.step_count_in_interval + for item, value in self.train_loss_accu.items() + } + + task_losses = self.train_loss_accu.get(_task_key, {}) + step_count = self.step_count_per_task.get(_task_key, 0) + if step_count == 0: + return dict.fromkeys(task_losses, float("nan")) + return { + item: value / step_count + for item, value in task_losses.items() + } else: def log_loss_train( @@ -1709,7 +1703,31 @@ def log_loss_valid(_task_key: str = "Default") -> dict: valid_results = {_key: {} for _key in self.model_keys} if self.disp_avg: # For multi-task, use accumulated average loss for all tasks + def initialize_task_loss_accumulator( + _task_key: str, + ) -> None: + if self.train_loss_accu[_task_key]: + return + self.optimizer.zero_grad(set_to_none=True) + task_input, task_label, _ = self.get_data( + is_train=True, task_key=_task_key + ) + if not task_input: + return + _, _, task_more_loss = self.wrapper( + **task_input, + cur_lr=pref_lr, + label=task_label, + task_key=_task_key, + ) + self.train_loss_accu[_task_key] = { + item: 0.0 + for item in task_more_loss + if "l2_" not in item + } + for _key in self.model_keys: + initialize_task_loss_accumulator(_key) train_results[_key] = log_loss_train( loss, more_loss, _task_key=_key ) @@ -1732,25 +1750,30 @@ def log_loss_valid(_task_key: str = "Default") -> dict: train_results[_key] = log_loss_train( loss, more_loss, _task_key=_key ) - valid_results[_key] = log_loss_valid(_task_key=_key) - if self.rank == 0: + for _key in self.model_keys: + valid_results[_key] = log_loss_valid(_task_key=_key) + if self.rank == 0: + log.info( + format_training_message_per_task( + batch=display_step_id, + task_name=_key + "_trn", + rmse=train_results[_key], + learning_rate=cur_lr, + check_total_rmse_nan=not ( + self.disp_avg + and self.step_count_per_task.get(_key, 0) == 0 + ), + ) + ) + if valid_results[_key]: log.info( format_training_message_per_task( batch=display_step_id, - task_name=_key + "_trn", - rmse=train_results[_key], - learning_rate=cur_lr, + task_name=_key + "_val", + rmse=valid_results[_key], + learning_rate=None, ) ) - if valid_results[_key]: - log.info( - format_training_message_per_task( - batch=display_step_id, - task_name=_key + "_val", - rmse=valid_results[_key], - learning_rate=None, - ) - ) self.wrapper.train() if self.disp_avg: diff --git a/source/tests/pt/test_multitask.py b/source/tests/pt/test_multitask.py index 560d89ed56..7937516ea8 100644 --- a/source/tests/pt/test_multitask.py +++ b/source/tests/pt/test_multitask.py @@ -9,6 +9,9 @@ from pathlib import ( Path, ) +from unittest.mock import ( + patch, +) import torch @@ -252,6 +255,34 @@ def setUp(self) -> None: self.config["model"] ) + def test_disp_avg_handles_unsampled_interval(self) -> None: + config = deepcopy(self.config) + config["training"]["numb_steps"] = 2 + config["training"]["disp_freq"] = 1 + config["training"]["disp_avg"] = True + config = update_deepmd_input(config, warning=True) + config = normalize(config, multi_task=True) + trainer = get_trainer(config, shared_links=self.shared_links) + + with patch( + "deepmd.pt.train.training.dp_random.choice", + side_effect=[0, 1], + ): + trainer.run() + + with open("lcurve.out") as f: + lines = f.readlines() + header_lines = [line.split() for line in lines if line.startswith("#")] + data_lines = [line.split() for line in lines if not line.startswith("#")] + self.assertTrue(header_lines) + header_columns = header_lines[0][1:] + self.assertTrue(any("_val_" in column for column in header_columns)) + for columns in data_lines: + self.assertEqual(len(columns), len(header_columns)) + displayed_steps = [int(columns[0]) for columns in data_lines] + self.assertEqual(displayed_steps, [1, 2]) + self.assertIn("nan", data_lines[1]) + def tearDown(self) -> None: MultiTaskTrainTest.tearDown(self)