已合并
[bugfix] fix pp save and load ckpt #8429
[bugfix] fix pp save and load ckpt #8429
已合并
yule100创建于 6月17日
共 8 个文件变更+86-41
@@ -31,6 +31,12 @@ from mindspore.nn.optim.optimizer import Optimizer
31from mindspore.communication import comm_func31from mindspore.communication import comm_func
32from mindspore import save_checkpoint as ms_save_checkpoint32from mindspore import save_checkpoint as ms_save_checkpoint
33 33 
34+try:
Sunshine_Youngster

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

likedislike
yule100
6月18日 评论:
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+ 
34from mindformers.checkpoint.layout_adapter import LayoutAdapter40from mindformers.checkpoint.layout_adapter import LayoutAdapter
35from mindformers.tools.logger import logger41from mindformers.tools.logger import logger
36from mindformers.checkpoint.reshard import ReshardLoader42from mindformers.checkpoint.reshard import ReshardLoader
@@ -264,7 +270,7 @@ class AsyncSaveManager:
264 return ten[0] == 0270 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."""
16import os16import os
17from collections import defaultdict17from collections import defaultdict
18-from typing import Callable18+from typing import Callable, List
19 19 
20from mindspore import save_checkpoint20from mindspore import save_checkpoint
21from mindspore.nn import Cell21from mindspore.nn import Cell
@@ -380,7 +380,7 @@ def distribute_shards(shard_coverage, shard_sizes, total_ranks):
380 return shard_assignment380 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, Tuple16+from typing import Dict, List, Union
17 17 
18import mindspore as ms18import mindspore as ms
19from mindspore import Parameter19from mindspore import Parameter
@@ -54,7 +54,7 @@ class LayoutAdapter:
54 return get_context("mode") == ms.context.PYNATIVE_MODE54 return get_context("mode") == ms.context.PYNATIVE_MODE
55 55 
56 @staticmethod56 @staticmethod
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_name205+ param_value = Parameter([])
198- param_value.requires_grad = param.requires_grad206+ param_value.data = param.to_local()
199- state_dict[param.name] = param_value207+ param_value.name = param_name
200- else:208+ param_value.requires_grad = param.requires_grad
201- state_dict[param.name] = param209+ state_dict[param.name] = param_value
210+ else:
211+ state_dict[param.name] = param
202 212 
203 return state_dict213 return state_dict
@@ -420,7 +420,7 @@ def get_sharded_tensor_from_cell(
420 420 
421 421 
422def get_all_sharded_tensor(422def get_all_sharded_tensor(
423- network: Cell,423+ network: Union[Cell, List[Cell]],
424 filter_func: Callable[[str], bool] = None424 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."""
16import os16import 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+ 
17from mindformers.pynative.callback.callback import TrainerCallback24from mindformers.pynative.callback.callback import TrainerCallback
18from mindformers.tools.logger import logger25from mindformers.tools.logger import logger
19from mindformers.tools.utils import get_real_group_size26from 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:
Sunshine_Youngster
已过期

if len(self.sharded_tensor_metas) == 0 and get_real_group_size() > 1:

likedislike
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 None173 ) 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 None180 ) 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 of128 # Update auxiliary-loss-free expert_bias on every step, regardless of
124 # log interval or loss availability. Non-last PP stages return loss=None129 # 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 = None133 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 return144 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 loss146 # 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 return66 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- return72+ _reset_max_attention_logit(m)
73+ return
P
Ppengjingyou6月18日

【功能】循环内return导致跳过了后续model的数据监控,return改为continue

likedislike
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 @staticmethod81 @staticmethod
81 def _fmt(v):82 def _fmt(v):
@@ -530,7 +530,7 @@ class Trainer:
530 # Create handler with complete list530 # 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,