已合并
[bugfix] fix pp save and load ckpt #8429
yule100创建于 6月17日
[bugfix] fix pp save and load ckpt #8429
已合并
共 8 个文件变更+86-41
| @@ -31,6 +31,12 @@ from mindspore.nn.optim.optimizer import Optimizer | |||
| 31 | from mindspore.communication import comm_func | 31 | from mindspore.communication import comm_func |
| 32 | from mindspore import save_checkpoint as ms_save_checkpoint | 32 | from mindspore import save_checkpoint as ms_save_checkpoint |
| 33 | 33 | ||
| 34 | +try: | ||
| 35 | + from hyper_parallel.core.distributed_checkpoint import get_global_layout | ||
| 36 | +except ImportError as e: | ||
| 37 | + get_global_layout = None | ||
| 38 | + logger.warning(f"Import get_global_layout failed: {e}.") | ||
| 39 | + | ||
| 34 | from mindformers.checkpoint.layout_adapter import LayoutAdapter | 40 | from mindformers.checkpoint.layout_adapter import LayoutAdapter |
| 35 | from mindformers.tools.logger import logger | 41 | from mindformers.tools.logger import logger |
| 36 | from mindformers.checkpoint.reshard import ReshardLoader | 42 | from mindformers.checkpoint.reshard import ReshardLoader |
| @@ -264,7 +270,7 @@ class AsyncSaveManager: | |||
| 264 | return ten[0] == 0 | 270 | return ten[0] == 0 |
| 265 | 271 | ||
| 266 | 272 | ||
| 267 | -def save_checkpoint(iteration: int, network: Cell, optimizer: Optimizer = None, | 273 | +def save_checkpoint(iteration: int, network: Union[Cell, List[Cell]], optimizer: Optimizer = None, |
| 268 | async_save_manager: AsyncSaveManager = None, common_info: CommonInfo = None, | 274 | async_save_manager: AsyncSaveManager = None, common_info: CommonInfo = None, |
| 269 | keep_max_num: int = 5, user_prefix: str = None, save_checkpoint_path: str = None, | 275 | keep_max_num: int = 5, user_prefix: str = None, save_checkpoint_path: str = None, |
| 270 | sharded_tensor_metas: Dict = None, remove_redundancy: bool = False, | 276 | sharded_tensor_metas: Dict = None, remove_redundancy: bool = False, |
| @@ -349,7 +355,19 @@ def save_checkpoint(iteration: int, network: Cell, optimizer: Optimizer = None, | |||
| 349 | 355 | ||
| 350 | # Save model weight. | 356 | # Save model weight. |
| 351 | logger.info("....... Start to save model weight .......") | 357 | logger.info("....... Start to save model weight .......") |
| 352 | - model_keys = network.parameters_dict().keys() | 358 | + if LayoutAdapter.is_pynative_mode() and get_real_group_size() > 1: |
| 359 | + # Get global model keys. | ||
| 360 | + network = network if isinstance(network, list) else [network] | ||
| 361 | + model_keys = set() | ||
| 362 | + for net in network: | ||
| 363 | + global_layout_dict = get_global_layout(net) | ||
| 364 | + for _, val in global_layout_dict.items(): | ||
| 365 | + model_keys.update(val.keys()) | ||
| 366 | + elif LayoutAdapter.is_pynative_mode(): | ||
| 367 | + network = network if isinstance(network, list) else [network] | ||
| 368 | + model_keys = network[0].parameters_dict().keys() | ||
| 369 | + else: | ||
| 370 | + model_keys = network.parameters_dict().keys() | ||
| 353 | start_save_ckpt_time = time() | 371 | start_save_ckpt_time = time() |
| 354 | 372 | ||
| 355 | if remove_redundancy and sharded_tensor_metas is not None: | 373 | if remove_redundancy and sharded_tensor_metas is not None: |
| @@ -15,7 +15,7 @@ | |||
| 15 | """save / load parallelization strategy.""" | 15 | """save / load parallelization strategy.""" |
| 16 | import os | 16 | import os |
| 17 | from collections import defaultdict | 17 | from collections import defaultdict |
| 18 | -from typing import Callable | 18 | +from typing import Callable, List |
| 19 | 19 | ||
| 20 | from mindspore import save_checkpoint | 20 | from mindspore import save_checkpoint |
| 21 | from mindspore.nn import Cell | 21 | from mindspore.nn import Cell |
| @@ -380,7 +380,7 @@ def distribute_shards(shard_coverage, shard_sizes, total_ranks): | |||
| 380 | return shard_assignment | 380 | return shard_assignment |
| 381 | 381 | ||
| 382 | 382 | ||
| 383 | -def apply_balance_shard_strategy(network: Cell, filter_func: Callable[[str], bool] = None): | 383 | +def apply_balance_shard_strategy(network: List[Cell], filter_func: Callable[[str], bool] = None): |
| 384 | """ | 384 | """ |
| 385 | Distributes and balances sharded tensor storage across ranks in a parallel group, | 385 | Distributes and balances sharded tensor storage across ranks in a parallel group, |
| 386 | generating rank-specific shard assignments. | 386 | generating rank-specific shard assignments. |
| @@ -13,7 +13,7 @@ | |||
| 13 | # limitations under the License. | 13 | # limitations under the License. |
| 14 | # ============================================================================ | 14 | # ============================================================================ |
| 15 | """Load/save checkpoint APIs for distributed parallel layout management.""" | 15 | """Load/save checkpoint APIs for distributed parallel layout management.""" |
| 16 | -from typing import Dict, Tuple | 16 | +from typing import Dict, List, Union |
| 17 | 17 | ||
| 18 | import mindspore as ms | 18 | import mindspore as ms |
| 19 | from mindspore import Parameter | 19 | from mindspore import Parameter |
| @@ -54,7 +54,7 @@ class LayoutAdapter: | |||
| 54 | return get_context("mode") == ms.context.PYNATIVE_MODE | 54 | return get_context("mode") == ms.context.PYNATIVE_MODE |
| 55 | 55 | ||
| 56 | 56 | ||
| 57 | - def get_all_layouts(network: Cell) -> Dict[int, Dict[str, list]]: | 57 | + def get_all_layouts(network: Union[Cell, List[Cell]]) -> Dict[int, Dict[str, list]]: |
| 58 | """ | 58 | """ |
| 59 | Retrieve distributed parallel layout information for all ranks in the network. | 59 | Retrieve distributed parallel layout information for all ranks in the network. |
| 60 | 60 | ||
| @@ -124,7 +124,15 @@ class LayoutAdapter: | |||
| 124 | """ | 124 | """ |
| 125 | if get_global_layout is None: | 125 | if get_global_layout is None: |
| 126 | raise ImportError("hyper_parallel is required for PyNative mode. Please install it.") | 126 | raise ImportError("hyper_parallel is required for PyNative mode. Please install it.") |
| 127 | - global_layout_dict = get_global_layout(network) | 127 | + global_layout_dict = {} |
| 128 | + network = network if isinstance(network, list) else [network] | ||
| 129 | + for net in network: | ||
| 130 | + layout_dict = get_global_layout(net) | ||
| 131 | + for rank_id, metas in layout_dict.items(): | ||
| 132 | + if rank_id in global_layout_dict: | ||
| 133 | + global_layout_dict[rank_id].update(metas) | ||
| 134 | + else: | ||
| 135 | + global_layout_dict[rank_id] = metas | ||
| 128 | 136 | ||
| 129 | if not global_layout_dict: | 137 | if not global_layout_dict: |
| 130 | return {} | 138 | return {} |
| @@ -190,14 +198,16 @@ class LayoutAdapter: | |||
| 190 | if DTensor is None: | 198 | if DTensor is None: |
| 191 | raise ImportError("DTensor is required for PyNative mode. Please install it.") | 199 | raise ImportError("DTensor is required for PyNative mode. Please install it.") |
| 192 | state_dict = {} | 200 | state_dict = {} |
| 193 | - for param_name, param in network.parameters_dict().items(): | 201 | + network = network if isinstance(network, list) else [network] |
| 194 | - if isinstance(param, DTensor): | 202 | + for net in network: |
| 195 | - param_value = Parameter([]) | 203 | + for param_name, param in net.parameters_dict().items(): |
| 196 | - param_value.data = param.to_local() | 204 | + if isinstance(param, DTensor): |
| 197 | - param_value.name = param_name | 205 | + param_value = Parameter([]) |
| 198 | - param_value.requires_grad = param.requires_grad | 206 | + param_value.data = param.to_local() |
| 199 | - state_dict[param.name] = param_value | 207 | + param_value.name = param_name |
| 200 | - else: | 208 | + param_value.requires_grad = param.requires_grad |
| 201 | - state_dict[param.name] = param | 209 | + state_dict[param.name] = param_value |
| 210 | + else: | ||
| 211 | + state_dict[param.name] = param | ||
| 202 | 212 | ||
| 203 | return state_dict | 213 | return state_dict |
| @@ -420,7 +420,7 @@ def get_sharded_tensor_from_cell( | |||
| 420 | 420 | ||
| 421 | 421 | ||
| 422 | def get_all_sharded_tensor( | 422 | def get_all_sharded_tensor( |
| 423 | - network: Cell, | 423 | + network: Union[Cell, List[Cell]], |
| 424 | filter_func: Callable[[str], bool] = None | 424 | filter_func: Callable[[str], bool] = None |
| 425 | ) -> Dict[int, Dict[str, ShardedTensor]]: | 425 | ) -> Dict[int, Dict[str, ShardedTensor]]: |
| 426 | """ | 426 | """ |
| @@ -14,6 +14,13 @@ | |||
| 14 | # ============================================================================ | 14 | # ============================================================================ |
| 15 | """Checkpoint callback for saving model checkpoints during training.""" | 15 | """Checkpoint callback for saving model checkpoints during training.""" |
| 16 | import os | 16 | import os |
| 17 | + | ||
| 18 | +try: | ||
| 19 | + from hyper_parallel.core.distributed_checkpoint import get_global_layout | ||
| 20 | +except ImportError as e: | ||
| 21 | + get_global_layout = None | ||
| 22 | + logger.warning(f"Import get_global_layout failed: {e}.") | ||
| 23 | + | ||
| 17 | from mindformers.pynative.callback.callback import TrainerCallback | 24 | from mindformers.pynative.callback.callback import TrainerCallback |
| 18 | from mindformers.tools.logger import logger | 25 | from mindformers.tools.logger import logger |
| 19 | from mindformers.tools.utils import get_real_group_size | 26 | from mindformers.tools.utils import get_real_group_size |
| @@ -153,11 +160,18 @@ class CheckpointCallback(TrainerCallback): | |||
| 153 | common_info = self._create_common_info(state) | 160 | common_info = self._create_common_info(state) |
| 154 | 161 | ||
| 155 | if self.sharded_tensor_metas is None and get_real_group_size() > 1: | 162 | if self.sharded_tensor_metas is None and get_real_group_size() > 1: |
已过期
![]() ![]() | |||
| 163 | + # Get global model keys. | ||
| 164 | + model_keys = set() | ||
| 165 | + for net in model: | ||
| 166 | + global_layout_dict = get_global_layout(net) | ||
| 167 | + for _, val in global_layout_dict.items(): | ||
| 168 | + model_keys.update(val.keys()) | ||
| 169 | + | ||
| 156 | self.sharded_tensor_metas = get_all_sharded_tensor( | 170 | self.sharded_tensor_metas = get_all_sharded_tensor( |
| 157 | network=model, | 171 | network=model, |
| 158 | - filter_func=(lambda x: x in list( | 172 | + filter_func=(lambda x: x in list(model_keys)) if self.no_save_optim else None |
| 159 | - model.parameters_dict().keys())) if self.no_save_optim else None | ||
| 160 | ) if get_real_group_size() > 1 else None | 173 | ) if get_real_group_size() > 1 else None |
| 174 | + | ||
| 161 | if self.opt_sharded_tensor_metas is None and get_real_group_size() > 1: | 175 | if self.opt_sharded_tensor_metas is None and get_real_group_size() > 1: |
| 162 | self.opt_sharded_tensor_metas = get_all_sharded_tensor( | 176 | self.opt_sharded_tensor_metas = get_all_sharded_tensor( |
| 163 | network=optimizer, | 177 | network=optimizer, |
| @@ -166,7 +180,7 @@ class CheckpointCallback(TrainerCallback): | |||
| 166 | ) if get_real_group_size() > 1 else None | 180 | ) if get_real_group_size() > 1 else None |
| 167 | 181 | ||
| 168 | if self.sharded_tensor_metas is not None and self.opt_sharded_tensor_metas is not None: | 182 | if self.sharded_tensor_metas is not None and self.opt_sharded_tensor_metas is not None: |
| 169 | - for rank_id, _ in self.sharded_tensor_metas.items(): | 183 | + for rank_id in self.sharded_tensor_metas: |
| 170 | self.sharded_tensor_metas[rank_id].update(self.opt_sharded_tensor_metas[rank_id]) | 184 | self.sharded_tensor_metas[rank_id].update(self.opt_sharded_tensor_metas[rank_id]) |
| 171 | 185 | ||
| 172 | try: | 186 | try: |
| @@ -119,6 +119,11 @@ class LossCallback(TrainerCallback): | |||
| 119 | model = kwargs.get("model") | 119 | model = kwargs.get("model") |
| 120 | metric_group = kwargs.get("metric_reduce_group") | 120 | metric_group = kwargs.get("metric_reduce_group") |
| 121 | metric_group_size = kwargs.get("metric_reduce_group_size") | 121 | metric_group_size = kwargs.get("metric_reduce_group_size") |
| 122 | + grad_norm = kwargs.get("grad_norm") | ||
| 123 | + | ||
| 124 | + # Calculate the time cost for the current step in milliseconds | ||
| 125 | + cur_time = time.time() | ||
| 126 | + step_time_cost = int((cur_time - self.step_time) * 1000) | ||
| 122 | 127 | ||
| 123 | # Update auxiliary-loss-free expert_bias on every step, regardless of | 128 | # Update auxiliary-loss-free expert_bias on every step, regardless of |
| 124 | # log interval or loss availability. Non-last PP stages return loss=None | 129 | # log interval or loss availability. Non-last PP stages return loss=None |
| @@ -127,20 +132,17 @@ class LossCallback(TrainerCallback): | |||
| 127 | # with the rest of the pipeline. | 132 | # with the rest of the pipeline. |
| 128 | model_config = None | 133 | model_config = None |
| 129 | if model is not None: | 134 | if model is not None: |
| 130 | - model_config = deepcopy(model.get_gpt_transformer_config()) | 135 | + model = model if isinstance(model, list) else [model] |
| 131 | - if getattr(model_config, "moe_router_enable_expert_bias", False): | 136 | + for m in model: |
| 132 | - _update_expert_bias(model, metric_group, metric_group_size) | 137 | + cfg = m.get_gpt_transformer_config() |
| 138 | + if getattr(cfg, "moe_router_enable_expert_bias", False): | ||
| 139 | + _update_expert_bias(m, metric_group, metric_group_size) | ||
| 140 | + model_config = deepcopy(m.get_gpt_transformer_config()) | ||
| 141 | + reset_model_temporary_tensors(model_config, m) | ||
| 142 | + | ||
| 133 | if loss is None or state.global_step % self.log_interval != 0: | 143 | if loss is None or state.global_step % self.log_interval != 0: |
| 134 | return | 144 | return |
| 135 | 145 | ||
| 136 | - grad_norm = kwargs.get("grad_norm") | ||
| 137 | - | ||
| 138 | - # Calculate the time cost for the current step in milliseconds | ||
| 139 | - cur_time = time.time() | ||
| 140 | - step_time_cost = int((cur_time - self.step_time) * 1000) | ||
| 141 | - | ||
| 142 | - reset_model_temporary_tensors(model_config, model) | ||
| 143 | - | ||
| 144 | # process aux loss | 146 | # process aux loss |
| 145 | load_balancing_loss = track_moe_metrics( | 147 | load_balancing_loss = track_moe_metrics( |
| 146 | loss_scale=model_config.moe_aux_loss_coeff, | 148 | loss_scale=model_config.moe_aux_loss_coeff, |
| @@ -65,17 +65,18 @@ class MaxLogitsMonitor(TrainerCallback): | |||
| 65 | _reset_max_attention_logit(model) | 65 | _reset_max_attention_logit(model) |
| 66 | return | 66 | return |
| 67 | 67 | ||
| 68 | - # 1) collect per-layer Parameter values. | 68 | + for m in model: |
| 69 | - params = model.get_max_attention_logit() | 69 | + # 1) collect per-layer Parameter values. |
| 70 | - if not params: | 70 | + params = m.get_max_attention_logit() |
| 71 | - _reset_max_attention_logit(model) | 71 | + if not params: |
| 72 | - return | 72 | + _reset_max_attention_logit(m) |
| 73 | + return | ||
P 【功能】循环内return导致跳过了后续model的数据监控,return改为continue ![]() ![]() | |||
| 73 | 74 | ||
| 74 | - # 2) dump. | 75 | + # 2) dump. |
| 75 | - self._dump(params, state) | 76 | + self._dump(params, state) |
| 76 | 77 | ||
| 77 | - # 3) reset for the next step. | 78 | + # 3) reset for the next step. |
| 78 | - _reset_max_attention_logit(model) | 79 | + _reset_max_attention_logit(m) |
| 79 | 80 | ||
| 80 | 81 | ||
| 81 | def _fmt(v): | 82 | def _fmt(v): |
| @@ -530,7 +530,7 @@ class Trainer: | |||
| 530 | # Create handler with complete list | 530 | # Create handler with complete list |
| 531 | cb_handler = CallbackHandler( | 531 | cb_handler = CallbackHandler( |
| 532 | callbacks=callback_list, | 532 | callbacks=callback_list, |
| 533 | - model=self.model[0], | 533 | + model=self.model, |
| 534 | train_dataset=self.train_dataset, | 534 | train_dataset=self.train_dataset, |
| 535 | eval_dataset=self.eval_dataset, | 535 | eval_dataset=self.eval_dataset, |
| 536 | optimizer=self.optimizer, | 536 | optimizer=self.optimizer, |


这里判断是动态图再import吧。静态图不依赖Hyper能力,就不交叉了。