| @@ -9,6 +9,7 @@ from torch import _C | |||
| 9 | from torch_npu.npu import _lazy_call, device as device_ctx_manager | 9 | from torch_npu.npu import _lazy_call, device as device_ctx_manager |
| 10 | from megatron.core.optimizer.cpu_offloading import HybridDeviceOptimizer | 10 | from megatron.core.optimizer.cpu_offloading import HybridDeviceOptimizer |
| 11 | from megatron.core.optimizer.distrib_optimizer import HAVE_APEX_OR_TE | 11 | from megatron.core.optimizer.distrib_optimizer import HAVE_APEX_OR_TE |
| 12 | +from mindspeed.core.optimizer.utils import _to_step_int | ||
| 12 | from mindspeed.core.tensor_parallel.tp_2d.group_api_2d import TPYCollectiveComm | 13 | from mindspeed.core.tensor_parallel.tp_2d.group_api_2d import TPYCollectiveComm |
| 13 | from mindspeed.core.tensor_parallel.tp_2d.layernorm_2d import LayerNorm2D | 14 | from mindspeed.core.tensor_parallel.tp_2d.layernorm_2d import LayerNorm2D |
| 14 | from mindspeed.core.tensor_parallel.tp_2d.rms_norm_2d import RMSNorm2D | 15 | from mindspeed.core.tensor_parallel.tp_2d.rms_norm_2d import RMSNorm2D |
| @@ -54,9 +55,12 @@ def _compile_dependencies(): | |||
| 54 | start_time = time.time() | 55 | start_time = time.time() |
| 55 | print('> compiling dataset index builder ...') | 56 | print('> compiling dataset index builder ...') |
| 56 | from megatron.core.datasets.utils import compile_helpers | 57 | from megatron.core.datasets.utils import compile_helpers |
| 58 | + | ||
| 57 | compile_helpers() | 59 | compile_helpers() |
| 58 | - print('>>> done with dataset index builder. Compilation time: {:.3f} ' | 60 | + print( |
| 59 | - 'seconds'.format(time.time() - start_time), flush=True) | 61 | + '>>> done with dataset index builder. Compilation time: {:.3f} seconds'.format(time.time() - start_time), |
| 62 | + flush=True, | ||
| 63 | + ) | ||
| 60 | 64 | ||
| 61 | 65 | ||
| 62 | def add_layer_norm_sp_support(config, instance): | 66 | def add_layer_norm_sp_support(config, instance): |
| @@ -69,9 +73,7 @@ def add_layer_norm_sp_support(config, instance): | |||
| 69 | setattr(instance, 'persist_layer_norm', persist_layer_norm) | 73 | setattr(instance, 'persist_layer_norm', persist_layer_norm) |
| 70 | 74 | ||
| 71 | 75 | ||
| 72 | - | ||
| 73 | class PTNorm: | 76 | class PTNorm: |
| 74 | - | ||
| 75 | def __new__(cls, config, hidden_size: int, eps: float = 1e-5): | 77 | def __new__(cls, config, hidden_size: int, eps: float = 1e-5): |
| 76 | if config.normalization == "LayerNorm": | 78 | if config.normalization == "LayerNorm": |
| 77 | if getattr(config, "tp_2d", False): | 79 | if getattr(config, "tp_2d", False): |
| @@ -84,6 +86,7 @@ class PTNorm: | |||
| 84 | try: | 86 | try: |
| 85 | # using apex implementation | 87 | # using apex implementation |
| 86 | from megatron.core.fusions.fused_layer_norm import FusedLayerNorm | 88 | from megatron.core.fusions.fused_layer_norm import FusedLayerNorm |
| 89 | + | ||
| 87 | instance = FusedLayerNorm(config=config, hidden_size=hidden_size, eps=eps) | 90 | instance = FusedLayerNorm(config=config, hidden_size=hidden_size, eps=eps) |
| 88 | except ImportError: | 91 | except ImportError: |
| 89 | # using torch implementation | 92 | # using torch implementation |
| @@ -99,10 +102,11 @@ class PTNorm: | |||
| 99 | instance.use_fused_rmsnorm = False | 102 | instance.use_fused_rmsnorm = False |
| 100 | else: | 103 | else: |
| 101 | from mindspeed.core.fusions.fused_rms_norm import RMSNorm | 104 | from mindspeed.core.fusions.fused_rms_norm import RMSNorm |
| 105 | + | ||
| 102 | instance = RMSNorm(dim=hidden_size, eps=eps, sequence_parallel=config.sequence_parallel, config=config) | 106 | instance = RMSNorm(dim=hidden_size, eps=eps, sequence_parallel=config.sequence_parallel, config=config) |
| 103 | instance.config.use_fused_rmsnorm = True | 107 | instance.config.use_fused_rmsnorm = True |
| 104 | else: | 108 | else: |
| 105 | - raise Exception('Only LayerNorm and RMSNorm are curently supported') | 109 | + raise ValueError('Only LayerNorm and RMSNorm are curently supported') |
| 106 | 110 | ||
| 107 | return instance | 111 | return instance |
| 108 | 112 | ||
| @@ -120,6 +124,7 @@ def get_device_wrapper(func): | |||
| 120 | else: | 124 | else: |
| 121 | device = func(*args, **kwargs) | 125 | device = func(*args, **kwargs) |
| 122 | return device | 126 | return device |
| 127 | + | ||
| 123 | return wrapper | 128 | return wrapper |
| 124 | 129 | ||
| 125 | 130 | ||
| @@ -142,7 +147,8 @@ def preload_tensors(write_buckets, non_blocking=True): | |||
| 142 | for bucket in write_buckets: | 147 | for bucket in write_buckets: |
| 143 | file_name, storage_key, (bytes_data, tensor_data) = bucket | 148 | file_name, storage_key, (bytes_data, tensor_data) = bucket |
| 144 | tensor_data = [ | 149 | tensor_data = [ |
| 145 | - (item, tensor.to("cpu", non_blocking=False) if not tensor.is_cpu else tensor.clone()) for item, tensor in tensor_data | 150 | + (item, tensor.to("cpu", non_blocking=False) if not tensor.is_cpu else tensor.clone()) |
| 151 | + for item, tensor in tensor_data | ||
| 146 | ] | 152 | ] |
| 147 | result.append((file_name, storage_key, (bytes_data, tensor_data))) | 153 | result.append((file_name, storage_key, (bytes_data, tensor_data))) |
| 148 | if non_blocking: | 154 | if non_blocking: |
| @@ -212,9 +218,7 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 212 | elif f"pre_{key}" in param_group: | 218 | elif f"pre_{key}" in param_group: |
| 213 | key = f"pre_{key}" | 219 | key = f"pre_{key}" |
| 214 | else: | 220 | else: |
| 215 | - raise ValueError( | 221 | + raise ValueError(f"Key {key} (or pre_{key}) not found in param_group {param_group}.") |
| 216 | - f"Key {key} (or pre_{key}) not found in param_group {param_group}." | ||
| 217 | - ) | ||
| 218 | needed_groups.append(param_group[key]) | 222 | needed_groups.append(param_group[key]) |
| 219 | needed_groups = tuple(needed_groups) | 223 | needed_groups = tuple(needed_groups) |
| 220 | return needed_groups | 224 | return needed_groups |
| @@ -227,9 +231,7 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 227 | state_dict_param_groups = [] | 231 | state_dict_param_groups = [] |
| 228 | for inner_param_group in inner_state_dict["param_groups"]: | 232 | for inner_param_group in inner_state_dict["param_groups"]: |
| 229 | needed_groups = make_needed_groups(inner_param_group) | 233 | needed_groups = make_needed_groups(inner_param_group) |
| 230 | - state_dict_param_groups.append( | 234 | + state_dict_param_groups.append({**param_groups_map[needed_groups], "params": inner_param_group['params']}) |
| 231 | - {**param_groups_map[needed_groups], "params": inner_param_group['params']} | ||
| 232 | - ) | ||
| 233 | 235 | ||
| 234 | # Allocate or retrieve optimizer state (i.e., tensors). | 236 | # Allocate or retrieve optimizer state (i.e., tensors). |
| 235 | if len(self.optimizer.state) == 0: | 237 | if len(self.optimizer.state) == 0: |
| @@ -245,13 +247,10 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 245 | for gbuf_range_map_for_all_buckets in gbuf_range_maps.values(): | 247 | for gbuf_range_map_for_all_buckets in gbuf_range_maps.values(): |
| 246 | for gbuf_range_map in gbuf_range_map_for_all_buckets: | 248 | for gbuf_range_map in gbuf_range_map_for_all_buckets: |
| 247 | for model_param, param_range_map in gbuf_range_map["param_map"].items(): | 249 | for model_param, param_range_map in gbuf_range_map["param_map"].items(): |
| 248 | - | ||
| 249 | # Get parameter ordering information (see method docstring | 250 | # Get parameter ordering information (see method docstring |
| 250 | # for details). | 251 | # for details). |
| 251 | group_index, group_order = self.model_param_group_index_map[model_param] | 252 | group_index, group_order = self.model_param_group_index_map[model_param] |
| 252 | - state_order = inner_state_dict["param_groups"][group_index]["params"][ | 253 | + state_order = inner_state_dict["param_groups"][group_index]["params"][group_order] |
| 253 | - group_order | ||
| 254 | - ] | ||
| 255 | 254 | ||
| 256 | # Allocate dummy tensors. | 255 | # Allocate dummy tensors. |
| 257 | numel = len(param_range_map["gbuf_world"]) | 256 | numel = len(param_range_map["gbuf_world"]) |
| @@ -276,7 +275,7 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 276 | 275 | ||
| 277 | # Extract 'step', for non-Apex/TE support. | 276 | # Extract 'step', for non-Apex/TE support. |
| 278 | if not HAVE_APEX_OR_TE: | 277 | if not HAVE_APEX_OR_TE: |
| 279 | - steps = list(set([g["step"] for g in state_dict["optimizer"]["param_groups"]])) | 278 | + steps = list({_to_step_int(g["step"]) for g in state_dict["optimizer"]["param_groups"]}) |
| 280 | if len(steps) != 1: | 279 | if len(steps) != 1: |
| 281 | raise AssertionError(f"Expect exactly one kind of step, but detect {len(steps)} kinds of steps") | 280 | raise AssertionError(f"Expect exactly one kind of step, but detect {len(steps)} kinds of steps") |
| 282 | step = torch.tensor(steps[0], dtype=torch.float) | 281 | step = torch.tensor(steps[0], dtype=torch.float) |
| @@ -287,9 +286,7 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 287 | elif isinstance(self.optimizer, HybridDeviceOptimizer): | 286 | elif isinstance(self.optimizer, HybridDeviceOptimizer): |
| 288 | # Handle Torch AdamW special case, which, unlike FusedAdam, Torch AdamW | 287 | # Handle Torch AdamW special case, which, unlike FusedAdam, Torch AdamW |
| 289 | # has an extra optimizer state “step”. | 288 | # has an extra optimizer state “step”. |
| 290 | - steps = list( | 289 | + steps = list({_to_step_int(g["step"]) for g in state_dict["optimizer"]["param_groups"] if "step" in g}) |
| 291 | - set([g["step"] for g in state_dict["optimizer"]["param_groups"] if "step" in g]) | ||
| 292 | - ) | ||
| 293 | if len(steps) != 0: | 290 | if len(steps) != 0: |
| 294 | if len(steps) != 1: | 291 | if len(steps) != 1: |
| 295 | raise AssertionError(f"steps: {steps}") | 292 | raise AssertionError(f"steps: {steps}") |
| @@ -298,16 +295,12 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 298 | v["step"] = step.detach().clone() | 295 | v["step"] = step.detach().clone() |
| 299 | 296 | ||
| 300 | # Optimizer. | 297 | # Optimizer. |
| 301 | - self.optimizer.load_state_dict( | 298 | + self.optimizer.load_state_dict({"state": state_dict_state, "param_groups": state_dict_param_groups}) |
| 302 | - {"state": state_dict_state, "param_groups": state_dict_param_groups} | ||
| 303 | - ) | ||
| 304 | 299 | ||
| 305 | # Grad scaler. | 300 | # Grad scaler. |
| 306 | if 'grad_scaler' not in state_dict: | 301 | if 'grad_scaler' not in state_dict: |
| 307 | if self.config.fp16: | 302 | if self.config.fp16: |
| 308 | - logger.info( | 303 | + logger.info('***WARNING*** found an old checkpoint, will not load grad scaler ...') |
| 309 | - '***WARNING*** found an old checkpoint, will not ' 'load grad scaler ...' | ||
| 310 | - ) | ||
| 311 | else: | 304 | else: |
| 312 | if self.grad_scaler: | 305 | if self.grad_scaler: |
| 313 | self.grad_scaler.load_state_dict(state_dict['grad_scaler']) | 306 | self.grad_scaler.load_state_dict(state_dict['grad_scaler']) |
| @@ -321,14 +314,14 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 321 | if 'param_state' in state_dict: | 314 | if 'param_state' in state_dict: |
| 322 | if 'param_state_sharding_type' not in state_dict: | 315 | if 'param_state_sharding_type' not in state_dict: |
| 323 | raise AssertionError( | 316 | raise AssertionError( |
| 324 | - f"Could not find 'param_state_sharding_type' in state_dict." | 317 | + f"Could not find 'param_state_sharding_type' in state_dict.Current state_dict.key(): {state_dict.key()}" |
| 325 | - f"Current state_dict.key(): {state_dict.key()}") | 318 | + ) |
| 326 | param_state = state_dict['param_state'] | 319 | param_state = state_dict['param_state'] |
| 327 | sharding_type = state_dict['param_state_sharding_type'] | 320 | sharding_type = state_dict['param_state_sharding_type'] |
| 328 | if self.ddp_config.use_custom_fsdp: | 321 | if self.ddp_config.use_custom_fsdp: |
| 329 | if sharding_type != "fully_sharded_model_space": | 322 | if sharding_type != "fully_sharded_model_space": |
| 330 | raise AssertionError("Only fully sharded model space is supported") | 323 | raise AssertionError("Only fully sharded model space is supported") |
| 331 | - logger.info(f'Loading distributed optimizer sharded state of type {sharding_type}') | 324 | + logger.info('Loading distributed optimizer sharded state of type %s', sharding_type) |
| 332 | if sharding_type == 'dp_zero_gather_scatter': | 325 | if sharding_type == 'dp_zero_gather_scatter': |
| 333 | self.load_parameter_state_from_dp_zero(param_state) | 326 | self.load_parameter_state_from_dp_zero(param_state) |
| 334 | elif sharding_type == 'fully_sharded_bucket_space': | 327 | elif sharding_type == 'fully_sharded_bucket_space': |
| @@ -336,4 +329,4 @@ def dist_optim_load_state_dict(self, state_dict): | |||
| 336 | elif sharding_type == 'fully_sharded_model_space': | 329 | elif sharding_type == 'fully_sharded_model_space': |
| 337 | self.load_parameter_state_from_fs_model_space(param_state) | 330 | self.load_parameter_state_from_fs_model_space(param_state) |
| 338 | else: | 331 | else: |
| 339 | - raise NotImplementedError(f'Unknown sharding_type: {sharding_type}') | 332 | + raise NotImplementedError(f'Unknown sharding_type: {sharding_type}') |
| @@ -1,25 +1,41 @@ | |||
| 1 | from typing import List, Optional, Tuple, Union | 1 | from typing import List, Optional, Tuple, Union |
| 2 | import torch | 2 | import torch |
| 3 | -import torch_npu | ||
| 4 | from torch import Tensor | 3 | from torch import Tensor |
| 5 | from torch.optim.optimizer import Optimizer | 4 | from torch.optim.optimizer import Optimizer |
| 6 | from torch.optim.adamw import AdamW as TorchAdamW | 5 | from torch.optim.adamw import AdamW as TorchAdamW |
| 6 | +from mindspeed.core.optimizer.utils import _to_step_int | ||
| 7 | 7 | ||
| 8 | 8 | ||
| 9 | -def adamw(params: List[Tensor], | 9 | +def _get_step_tensor(optimizer, group): |
| 10 | - grads: List[Tensor], | 10 | + device = torch.npu.current_device() |
| 11 | - exp_avgs: List[Tensor], | 11 | + step_tensor_cache = getattr(optimizer, '_step_tensor_cache', None) |
| 12 | - exp_avg_sqs: List[Tensor], | 12 | + if step_tensor_cache is None: |
| 13 | - max_exp_avg_sqs: List[Tensor], | 13 | + step_tensor_cache = {} |
| 14 | - step_tensor: Tensor, | 14 | + optimizer._step_tensor_cache = step_tensor_cache |
| 15 | - *, | 15 | + step_tensor = step_tensor_cache.get(id(group)) |
| 16 | - amsgrad: bool, | 16 | + if step_tensor is None or step_tensor.device.index != device: |
| 17 | - beta1: float, | 17 | + step_tensor = torch.empty((), dtype=torch.int64, device=device) |
| 18 | - beta2: float, | 18 | + step_tensor_cache[id(group)] = step_tensor |
| 19 | - lr: float, | 19 | + step_tensor.fill_(group['step']) |
| 20 | - weight_decay: float, | 20 | + return step_tensor |
| 21 | - eps: float, | 21 | + |
| 22 | - maximize: bool): | 22 | + |
| 23 | +def adamw( | ||
| 24 | + params: List[Tensor], | ||
| 25 | + grads: List[Tensor], | ||
| 26 | + exp_avgs: List[Tensor], | ||
| 27 | + exp_avg_sqs: List[Tensor], | ||
| 28 | + max_exp_avg_sqs: List[Tensor], | ||
| 29 | + step_tensor: Tensor, | ||
| 30 | + *, | ||
| 31 | + amsgrad: bool, | ||
| 32 | + beta1: float, | ||
| 33 | + beta2: float, | ||
| 34 | + lr: float, | ||
| 35 | + weight_decay: float, | ||
| 36 | + eps: float, | ||
| 37 | + maximize: bool, | ||
| 38 | +): | ||
| 23 | r"""Functional API that performs AdamW algorithm computation. | 39 | r"""Functional API that performs AdamW algorithm computation. |
| 24 | See :class:`~torch.optim.AdamW` for details. | 40 | See :class:`~torch.optim.AdamW` for details. |
| 25 | """ | 41 | """ |
| @@ -42,7 +58,7 @@ def adamw(params: List[Tensor], | |||
| 42 | beta2=beta2, | 58 | beta2=beta2, |
| 43 | weight_decay=weight_decay, | 59 | weight_decay=weight_decay, |
| 44 | eps=eps, | 60 | eps=eps, |
| 45 | - maximize=maximize | 61 | + maximize=maximize, |
| 46 | ) | 62 | ) |
| 47 | 63 | ||
| 48 | 64 | ||
| @@ -62,42 +78,50 @@ class FusedTorchAdamW(TorchAdamW): | |||
| 62 | differentiable: bool = False, | 78 | differentiable: bool = False, |
| 63 | fused: Optional[bool] = None, | 79 | fused: Optional[bool] = None, |
| 64 | ): | 80 | ): |
| 65 | - super().__init__(params, | 81 | + super().__init__( |
| 66 | - lr=lr, | 82 | + params, |
| 67 | - betas=betas, | 83 | + lr=lr, |
| 68 | - eps=eps, | 84 | + betas=betas, |
| 69 | - weight_decay=weight_decay, | 85 | + eps=eps, |
| 70 | - amsgrad=amsgrad, | 86 | + weight_decay=weight_decay, |
| 71 | - foreach=False, | 87 | + amsgrad=amsgrad, |
| 72 | - maximize=maximize, | 88 | + foreach=False, |
| 73 | - capturable=False, | 89 | + maximize=maximize, |
| 74 | - differentiable=False, | 90 | + capturable=False, |
| 75 | - fused=True,) | 91 | + differentiable=False, |
| 92 | + fused=True, | ||
| 93 | + ) | ||
| 76 | 94 | ||
| 77 | 95 | ||
| 78 | class AdamW(Optimizer): | 96 | class AdamW(Optimizer): |
| 79 | - def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, | 97 | + def __init__( |
| 80 | - weight_decay=1e-2, amsgrad=False, *, maximize: bool = False): | 98 | + self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=1e-2, amsgrad=False, *, maximize: bool = False |
| 81 | - if not 0.0 <= lr: | 99 | + ): |
| 100 | + if 0.0 > lr: | ||
| 82 | raise ValueError("Invalid learning rate: {}".format(lr)) | 101 | raise ValueError("Invalid learning rate: {}".format(lr)) |
| 83 | - if not 0.0 <= eps: | 102 | + if 0.0 > eps: |
| 84 | raise ValueError("Invalid epsilon value: {}".format(eps)) | 103 | raise ValueError("Invalid epsilon value: {}".format(eps)) |
| 85 | if not 0.0 <= betas[0] < 1.0: | 104 | if not 0.0 <= betas[0] < 1.0: |
| 86 | raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) | 105 | raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) |
| 87 | if not 0.0 <= betas[1] < 1.0: | 106 | if not 0.0 <= betas[1] < 1.0: |
| 88 | raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) | 107 | raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) |
| 89 | - if not 0.0 <= weight_decay: | 108 | + if 0.0 > weight_decay: |
| 90 | raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) | 109 | raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) |
| 91 | - defaults = dict(lr=lr, betas=betas, eps=eps, | 110 | + defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay, amsgrad=amsgrad, maximize=maximize) |
| 92 | - weight_decay=weight_decay, amsgrad=amsgrad, maximize=maximize) | 111 | + super().__init__(params, defaults) |
| 93 | - super(AdamW, self).__init__(params, defaults) | 112 | + self._step_tensor_cache = {} |
| 94 | 113 | ||
| 95 | def __setstate__(self, state): | 114 | def __setstate__(self, state): |
| 96 | - super(AdamW, self).__setstate__(state) | 115 | + super().__setstate__(state) |
| 97 | for group in self.param_groups: | 116 | for group in self.param_groups: |
| 98 | group.setdefault('amsgrad', False) | 117 | group.setdefault('amsgrad', False) |
| 99 | group.setdefault('maximize', False) | 118 | group.setdefault('maximize', False) |
| 100 | 119 | ||
| 120 | + def load_state_dict(self, state_dict): | ||
| 121 | + result = super().load_state_dict(state_dict) | ||
| 122 | + self._step_tensor_cache = {} | ||
| 123 | + return result | ||
| 124 | + | ||
| 101 | 125 | ||
| 102 | def step(self, closure=None): | 126 | def step(self, closure=None): |
| 103 | loss = None | 127 | loss = None |
| @@ -110,18 +134,15 @@ class AdamW(Optimizer): | |||
| 110 | grads = [] | 134 | grads = [] |
| 111 | exp_avgs = [] | 135 | exp_avgs = [] |
| 112 | exp_avg_sqs = [] | 136 | exp_avg_sqs = [] |
| 113 | - state_sums = [] | ||
| 114 | max_exp_avg_sqs = [] | 137 | max_exp_avg_sqs = [] |
| 115 | - state_steps = [] | ||
| 116 | amsgrad = group['amsgrad'] | 138 | amsgrad = group['amsgrad'] |
| 117 | beta1, beta2 = group['betas'] | 139 | beta1, beta2 = group['betas'] |
| 118 | 140 | ||
| 119 | if 'step' in group: | 141 | if 'step' in group: |
| 120 | - group['step'] += 1 | 142 | + group['step'] = _to_step_int(group['step']) + 1 |
| 121 | - if group['step'].is_cpu: | ||
| 122 | - group['step'] = group['step'].cuda() | ||
| 123 | else: | 143 | else: |
| 124 | - group['step'] = torch.tensor(1, dtype=torch.int64, device=torch.cuda.current_device()) | 144 | + group['step'] = 1 |
| 145 | + step_tensor = _get_step_tensor(self, group) | ||
| 125 | 146 | ||
| 126 | for p in group['params']: | 147 | for p in group['params']: |
| 127 | if p.grad is None: | 148 | if p.grad is None: |
| @@ -149,18 +170,20 @@ class AdamW(Optimizer): | |||
| 149 | if amsgrad: | 170 | if amsgrad: |
| 150 | max_exp_avg_sqs.append(state['max_exp_avg_sq']) | 171 | max_exp_avg_sqs.append(state['max_exp_avg_sq']) |
| 151 | 172 | ||
| 152 | - adamw(params_with_grad, | 173 | + adamw( |
| 153 | - grads, | 174 | + params_with_grad, |
| 154 | - exp_avgs, | 175 | + grads, |
| 155 | - exp_avg_sqs, | 176 | + exp_avgs, |
| 156 | - max_exp_avg_sqs, | 177 | + exp_avg_sqs, |
| 157 | - group['step'], | 178 | + max_exp_avg_sqs, |
| 158 | - amsgrad=amsgrad, | 179 | + step_tensor, |
| 159 | - beta1=beta1, | 180 | + amsgrad=amsgrad, |
| 160 | - beta2=beta2, | 181 | + beta1=beta1, |
| 161 | - lr=group['lr'], | 182 | + beta2=beta2, |
| 162 | - weight_decay=group['weight_decay'], | 183 | + lr=group['lr'], |
| 163 | - eps=group['eps'], | 184 | + weight_decay=group['weight_decay'], |
| 164 | - maximize=group['maximize']) | 185 | + eps=group['eps'], |
| 186 | + maximize=group['maximize'], | ||
| 187 | + ) | ||
| 165 | 188 | ||
| 166 | return loss | 189 | return loss |
| @@ -0,0 +1,16 @@ | |||
| 1 | +# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. All rights reserved. | ||
| 3 | + | ||
| 4 | +import torch | ||
| 5 | +from mindspeed.core.optimizer.utils import _distributed_group_rank | ||
| 6 | + | ||
| 7 | + | ||
| 8 | +def load_parameter_state(self, filename: str, *, update_legacy_format=False): | ||
| 9 | + """Load distributed optimizer parameter state through CPU for cross-device resume.""" | ||
| 10 | + if self.is_stub_optimizer: | ||
| 11 | + return | ||
| 12 | + state_dict = None | ||
| 13 | + if _distributed_group_rank(self.data_parallel_group) == 0: | ||
| 14 | + state_dict = torch.load(filename, map_location='cpu') | ||
| 15 | + | ||
| 16 | + self.load_parameter_state_from_dp_zero(state_dict, update_legacy_format=update_legacy_format) | ||
| @@ -0,0 +1,23 @@ | |||
| 1 | +# Copyright (c) 2024, NVIDIA CORPORATION. All rights reserved. | ||
| 2 | +# Copyright (c) 2026, Huawei Technologies Co., Ltd. All rights reserved. | ||
| 3 | + | ||
| 4 | +import torch | ||
| 5 | +from mindspeed.core.optimizer.utils import _distributed_group_rank | ||
| 6 | + | ||
| 7 | + | ||
| 8 | +def load_parameter_state(self, filename: str, *, update_legacy_format: bool = False): | ||
| 9 | + """Load chained distributed optimizer parameter states through CPU.""" | ||
| 10 | + if len(self.chained_optimizers) == 1: | ||
| 11 | + self.chained_optimizers[0].load_parameter_state(filename, update_legacy_format=update_legacy_format) | ||
| 12 | + return | ||
| 13 | + | ||
| 14 | + states = None | ||
| 15 | + for idx, optimizer in enumerate(self.chained_optimizers): | ||
| 16 | + if not hasattr(optimizer, 'load_parameter_state_from_dp_zero'): | ||
| 17 | + continue | ||
| 18 | + | ||
| 19 | + if _distributed_group_rank(optimizer.data_parallel_group) == 0 and states is None: | ||
| 20 | + states = torch.load(filename, map_location='cpu') | ||
| 21 | + | ||
| 22 | + state_dict = states[idx] if states else None | ||
🟡 Medium Priority 在 当某个 optimizer 不满足 此问题在上一轮审查中已提出,当前 diff 中未修复。虽然在实际使用中所有 chained_optimizers 通常都是 DistributedOptimizer 实例(均有 建议:建议引入一个独立的计数器 ![]() ![]() 不准确? | |||
| 23 | + optimizer.load_parameter_state_from_dp_zero(state_dict, update_legacy_format=update_legacy_format) | ||
| @@ -0,0 +1,13 @@ | |||
| 1 | +import torch | ||
| 2 | + | ||
| 3 | + | ||
| 4 | +def _to_step_int(step): | ||
| 5 | + if torch.is_tensor(step): | ||
| 6 | + return int(step.item()) | ||
| 7 | + return int(step) | ||
| 8 | + | ||
| 9 | + | ||
| 10 | +def _distributed_group_rank(group): | ||
| 11 | + if hasattr(group, 'rank') and callable(group.rank): | ||
| 12 | + return group.rank() | ||
| 13 | + return torch.distributed.get_rank(group) | ||
| @@ -12,30 +12,33 @@ class MegatronBasicFeature(MindSpeedFeature): | |||
| 12 | 12 | ||
| 13 | def validate_args(self, args): | 13 | def validate_args(self, args): |
| 14 | # Fix VPP when VPP_size=1 from megatron core_r0.14.0 (!3640). | 14 | # Fix VPP when VPP_size=1 from megatron core_r0.14.0 (!3640). |
| 15 | - if getattr(args, 'num_layers_per_virtual_pipeline_stage', None) is not None or getattr(args, 'num_virtual_stages_per_pipeline_rank', None) is not None: | 15 | + if ( |
| 16 | + getattr(args, 'num_layers_per_virtual_pipeline_stage', None) is not None | ||
| 17 | + or getattr(args, 'num_virtual_stages_per_pipeline_rank', None) is not None | ||
| 18 | + ): | ||
| 16 | if args.virtual_pipeline_model_parallel_size == 1 and not getattr(args, 'moe_fb_overlap', False): | 19 | if args.virtual_pipeline_model_parallel_size == 1 and not getattr(args, 'moe_fb_overlap', False): |
| 17 | args.virtual_pipeline_model_parallel_size = None | 20 | args.virtual_pipeline_model_parallel_size = None |
| 18 | args.overlap_p2p_comm = False | 21 | args.overlap_p2p_comm = False |
| 19 | - | 22 | + |
| 20 | - if (getattr(args, 'num_layers_per_virtual_pipeline_stage', None) is not None and | 23 | + if ( |
| 21 | - getattr(args, 'pipeline_model_parallel_size', None) is not None and | 24 | + getattr(args, 'num_layers_per_virtual_pipeline_stage', None) is not None |
| 22 | - args.num_layers_per_virtual_pipeline_stage * args.pipeline_model_parallel_size == args.num_layers): | 25 | + and getattr(args, 'pipeline_model_parallel_size', None) is not None |
| 26 | + and args.num_layers_per_virtual_pipeline_stage * args.pipeline_model_parallel_size == args.num_layers | ||
| 27 | + ): | ||
| 23 | raise ValueError( | 28 | raise ValueError( |
| 24 | 'num_layers_per_virtual_pipeline_stage * pipeline_model_parallel_size == num_layers, ' | 29 | 'num_layers_per_virtual_pipeline_stage * pipeline_model_parallel_size == num_layers, ' |
| 25 | 'please close --num-layers-per-virtual-pipeline-stage' | 30 | 'please close --num-layers-per-virtual-pipeline-stage' |
| 26 | ) | 31 | ) |
| 27 | - | 32 | + |
| 28 | if getattr(args, 'defer_embedding_wgrad_compute', False): | 33 | if getattr(args, 'defer_embedding_wgrad_compute', False): |
| 29 | raise AssertionError( | 34 | raise AssertionError( |
| 30 | '--defer_embedding_wgrad_compute, although exclusive to TE scenarios, is not yet supported.' | 35 | '--defer_embedding_wgrad_compute, although exclusive to TE scenarios, is not yet supported.' |
| 31 | ) | 36 | ) |
| 32 | 37 | ||
| 33 | def register_patches(self, patch_manager, args): | 38 | def register_patches(self, patch_manager, args): |
| 34 | - try: | 39 | + import importlib.util |
| 35 | - import megatron.training | 40 | + |
| 36 | - only_mcore = False | 41 | + only_mcore = importlib.util.find_spec("megatron.training") is None |
| 37 | - except ModuleNotFoundError: | ||
| 38 | - only_mcore = True | ||
| 39 | 42 | ||
| 40 | self.register_mcore_basic_patches(patch_manager, args) | 43 | self.register_mcore_basic_patches(patch_manager, args) |
| 41 | if not only_mcore: | 44 | if not only_mcore: |
| @@ -43,20 +46,36 @@ class MegatronBasicFeature(MindSpeedFeature): | |||
| 43 | 46 | ||
| 44 | def register_mcore_basic_patches(self, pm, args): | 47 | def register_mcore_basic_patches(self, pm, args): |
| 45 | # configuration patches | 48 | # configuration patches |
| 46 | - from mindspeed.core.megatron_basic.arguments_basic import (transformer_config_init_wrapper, | 49 | + from mindspeed.core.megatron_basic.arguments_basic import ( |
| 47 | - transformer_config_post_init_wrapper, | 50 | + transformer_config_init_wrapper, |
| 48 | - transformer_config_init_subclass) | 51 | + transformer_config_post_init_wrapper, |
| 49 | - pm.register_patch("megatron.core.transformer.transformer_config.TransformerConfig.__init__", transformer_config_init_wrapper) | 52 | + transformer_config_init_subclass, |
| 50 | - pm.register_patch("megatron.core.transformer.transformer_config.TransformerConfig.__init_subclass__", classmethod(transformer_config_init_subclass)) | 53 | + ) |
| 51 | - pm.register_patch("megatron.core.transformer.transformer_config.TransformerConfig.__post_init__", transformer_config_post_init_wrapper) | 54 | + |
| 52 | - pm.register_patch("megatron.core.transformer.transformer_config.MLATransformerConfig.__init__", transformer_config_init_wrapper) | 55 | + pm.register_patch( |
| 56 | + "megatron.core.transformer.transformer_config.TransformerConfig.__init__", transformer_config_init_wrapper | ||
| 57 | + ) | ||
| 58 | + pm.register_patch( | ||
| 59 | + "megatron.core.transformer.transformer_config.TransformerConfig.__init_subclass__", | ||
| 60 | + classmethod(transformer_config_init_subclass), | ||
| 61 | + ) | ||
| 62 | + pm.register_patch( | ||
| 63 | + "megatron.core.transformer.transformer_config.TransformerConfig.__post_init__", | ||
| 64 | + transformer_config_post_init_wrapper, | ||
| 65 | + ) | ||
| 66 | + pm.register_patch( | ||
| 67 | + "megatron.core.transformer.transformer_config.MLATransformerConfig.__init__", | ||
| 68 | + transformer_config_init_wrapper, | ||
| 69 | + ) | ||
| 53 | 70 | ||
| 54 | # initialization patches | 71 | # initialization patches |
| 55 | from mindspeed.core.megatron_basic.megatron_basic import _set_cuda_rng_state | 72 | from mindspeed.core.megatron_basic.megatron_basic import _set_cuda_rng_state |
| 73 | + | ||
| 56 | pm.register_patch('megatron.core.tensor_parallel.random._set_cuda_rng_state', _set_cuda_rng_state) | 74 | pm.register_patch('megatron.core.tensor_parallel.random._set_cuda_rng_state', _set_cuda_rng_state) |
| 57 | 75 | ||
| 58 | # norm patches | 76 | # norm patches |
| 59 | from mindspeed.core.megatron_basic.megatron_basic import PTNorm | 77 | from mindspeed.core.megatron_basic.megatron_basic import PTNorm |
| 78 | + | ||
| 60 | pm.register_patch('megatron.core.models.gpt.gpt_layer_specs.LNImpl', PTNorm) | 79 | pm.register_patch('megatron.core.models.gpt.gpt_layer_specs.LNImpl', PTNorm) |
| 61 | pm.register_patch('megatron.core.transformer.torch_norm.WrappedTorchNorm', PTNorm) | 80 | pm.register_patch('megatron.core.transformer.torch_norm.WrappedTorchNorm', PTNorm) |
| 62 | pm.register_patch('megatron.core.transformer.transformer_block.LayerNormImpl', PTNorm) | 81 | pm.register_patch('megatron.core.transformer.transformer_block.LayerNormImpl', PTNorm) |
| @@ -66,29 +85,62 @@ class MegatronBasicFeature(MindSpeedFeature): | |||
| 66 | from mindspeed.core.optimizer.fix_duplicate_allgather import start_param_sync | 85 | from mindspeed.core.optimizer.fix_duplicate_allgather import start_param_sync |
| 67 | from mindspeed.core.optimizer.fix_duplicate_allgather import step_with_ready_grads_distrib_opti_wrapper | 86 | from mindspeed.core.optimizer.fix_duplicate_allgather import step_with_ready_grads_distrib_opti_wrapper |
| 68 | from mindspeed.core.optimizer.fix_duplicate_allgather import get_megatron_optimizer_wrapper | 87 | from mindspeed.core.optimizer.fix_duplicate_allgather import get_megatron_optimizer_wrapper |
| 69 | - pm.register_patch('megatron.core.distributed.distributed_data_parallel.DistributedDataParallel.start_param_sync', start_param_sync) | 88 | + |
| 70 | - pm.register_patch('megatron.core.optimizer.distrib_optimizer.DistributedOptimizer.step_with_ready_grads', step_with_ready_grads_distrib_opti_wrapper) | 89 | + pm.register_patch( |
| 90 | + 'megatron.core.distributed.distributed_data_parallel.DistributedDataParallel.start_param_sync', | ||
| 91 | + start_param_sync, | ||
| 92 | + ) | ||
| 93 | + pm.register_patch( | ||
| 94 | + 'megatron.core.optimizer.distrib_optimizer.DistributedOptimizer.step_with_ready_grads', | ||
| 95 | + step_with_ready_grads_distrib_opti_wrapper, | ||
| 96 | + ) | ||
| 71 | pm.register_patch('megatron.core.optimizer.get_megatron_optimizer', get_megatron_optimizer_wrapper) | 97 | pm.register_patch('megatron.core.optimizer.get_megatron_optimizer', get_megatron_optimizer_wrapper) |
| 72 | 98 | ||
| 73 | # Currently, it is not supported to Cast shard fp32 main params to fp8 model params | 99 | # Currently, it is not supported to Cast shard fp32 main params to fp8 model params |
| 74 | from mindspeed.core.fp8_utils import quantize_param_shard | 100 | from mindspeed.core.fp8_utils import quantize_param_shard |
| 101 | + | ||
| 75 | pm.register_patch('megatron.core.fp8_utils.quantize_param_shard', quantize_param_shard) | 102 | pm.register_patch('megatron.core.fp8_utils.quantize_param_shard', quantize_param_shard) |
| 76 | 103 | ||
| 77 | # fix count_zeros in ChainedOptimizer for core_r0.12.1. | 104 | # fix count_zeros in ChainedOptimizer for core_r0.12.1. |
| 78 | from mindspeed.core.megatron_basic.count_zero_fix import step | 105 | from mindspeed.core.megatron_basic.count_zero_fix import step |
| 106 | + | ||
| 79 | pm.register_patch('megatron.core.optimizer.optimizer.ChainedOptimizer.step', step) | 107 | pm.register_patch('megatron.core.optimizer.optimizer.ChainedOptimizer.step', step) |
| 80 | 108 | ||
| 81 | # avoid async save | 109 | # avoid async save |
| 82 | from mindspeed.core.megatron_basic.megatron_basic import preload_tensors | 110 | from mindspeed.core.megatron_basic.megatron_basic import preload_tensors |
| 83 | - pm.register_patch('megatron.core.dist_checkpointing.strategies.filesystem_async.FileSystemWriterAsync.preload_tensors', preload_tensors) | 111 | + |
| 112 | + pm.register_patch( | ||
| 113 | + 'megatron.core.dist_checkpointing.strategies.filesystem_async.FileSystemWriterAsync.preload_tensors', | ||
| 114 | + preload_tensors, | ||
| 115 | + ) | ||
| 84 | 116 | ||
| 85 | # avoid incorrect weight_decay override in resume task | 117 | # avoid incorrect weight_decay override in resume task |
| 86 | from mindspeed.core.megatron_basic.megatron_basic import dist_optim_load_state_dict | 118 | from mindspeed.core.megatron_basic.megatron_basic import dist_optim_load_state_dict |
| 87 | - pm.register_patch('megatron.core.optimizer.distrib_optimizer.DistributedOptimizer.load_state_dict', dist_optim_load_state_dict) | 119 | + from mindspeed.core.optimizer.distrib_optimizer import ( |
| 120 | + load_parameter_state as distributed_optimizer_load_parameter_state, | ||
| 121 | + ) | ||
| 122 | + from mindspeed.core.optimizer.optimizer import load_parameter_state as chained_optimizer_load_parameter_state | ||
| 123 | + | ||
| 124 | + pm.register_patch( | ||
| 125 | + 'megatron.core.optimizer.distrib_optimizer.DistributedOptimizer.load_state_dict', dist_optim_load_state_dict | ||
| 126 | + ) | ||
| 127 | + pm.register_patch( | ||
| 128 | + 'megatron.core.optimizer.distrib_optimizer.DistributedOptimizer.load_parameter_state', | ||
| 129 | + distributed_optimizer_load_parameter_state, | ||
| 130 | + ) | ||
| 131 | + pm.register_patch( | ||
| 132 | + 'megatron.core.optimizer.optimizer.ChainedOptimizer.load_parameter_state', | ||
| 133 | + chained_optimizer_load_parameter_state, | ||
| 134 | + ) | ||
| 88 | 135 | ||
| 89 | def register_non_mcore_basic_patches(self, pm, args): | 136 | def register_non_mcore_basic_patches(self, pm, args): |
| 90 | # args parser patch | 137 | # args parser patch |
| 91 | - from mindspeed.core.megatron_basic.arguments_basic import parse_args_wrapper, validate_args_wrapper, print_args_wrapper | 138 | + from mindspeed.core.megatron_basic.arguments_basic import ( |
| 139 | + parse_args_wrapper, | ||
| 140 | + validate_args_wrapper, | ||
| 141 | + print_args_wrapper, | ||
| 142 | + ) | ||
| 143 | + | ||
| 92 | pm.register_patch('megatron.training.arguments.parse_args', parse_args_wrapper) | 144 | pm.register_patch('megatron.training.arguments.parse_args', parse_args_wrapper) |
| 93 | pm.register_patch('megatron.training.arguments.validate_args', validate_args_wrapper) | 145 | pm.register_patch('megatron.training.arguments.validate_args', validate_args_wrapper) |
| 94 | pm.register_patch('megatron.training.arguments._print_args', print_args_wrapper) | 146 | pm.register_patch('megatron.training.arguments._print_args', print_args_wrapper) |
| @@ -97,10 +149,10 @@ class MegatronBasicFeature(MindSpeedFeature): | |||
| 97 | 149 | ||
| 98 | # initialization patches | 150 | # initialization patches |
| 99 | from mindspeed.core.megatron_basic.megatron_basic import _compile_dependencies, get_device_wrapper | 151 | from mindspeed.core.megatron_basic.megatron_basic import _compile_dependencies, get_device_wrapper |
| 152 | + | ||
| 100 | pm.register_patch('megatron.training.initialize._compile_dependencies', _compile_dependencies) | 153 | pm.register_patch('megatron.training.initialize._compile_dependencies', _compile_dependencies) |
| 101 | pm.register_patch('megatron.training.dist_signal_handler.get_device', get_device_wrapper) | 154 | pm.register_patch('megatron.training.dist_signal_handler.get_device', get_device_wrapper) |
| 102 | 155 | ||
| 103 | from mindspeed.core.megatron_basic.megatron_basic import get_device_arch_version | 156 | from mindspeed.core.megatron_basic.megatron_basic import get_device_arch_version |
| 157 | + | ||
| 104 | pm.register_patch('megatron.training.utils.get_device_arch_version', get_device_arch_version) | 158 | pm.register_patch('megatron.training.utils.get_device_arch_version', get_device_arch_version) |
| 105 | - | ||
| 106 | - | ||
| @@ -1,25 +1,40 @@ | |||
| 1 | -from typing import List, Optional, Tuple, Union | 1 | +from typing import List |
| 2 | import torch | 2 | import torch |
| 3 | -import torch_npu | ||
| 4 | from torch import Tensor | 3 | from torch import Tensor |
| 5 | from torch.optim.optimizer import Optimizer | 4 | from torch.optim.optimizer import Optimizer |
| 6 | -from torch.optim.adamw import AdamW as TorchAdamW | 5 | +from mindspeed.core.optimizer.utils import _to_step_int |
| 7 | 6 | ||
| 8 | 7 | ||
| 9 | -def adamw(params: List[Tensor], | 8 | +def _get_step_tensor(optimizer, group): |
| 10 | - grads: List[Tensor], | 9 | + device = torch.npu.current_device() |
| 11 | - exp_avgs: List[Tensor], | 10 | + step_tensor_cache = getattr(optimizer, '_step_tensor_cache', None) |
| 12 | - exp_avg_sqs: List[Tensor], | 11 | + if step_tensor_cache is None: |
| 13 | - max_exp_avg_sqs: List[Tensor], | 12 | + step_tensor_cache = {} |
| 14 | - step_tensor: Tensor, | 13 | + optimizer._step_tensor_cache = step_tensor_cache |
| 15 | - *, | 14 | + step_tensor = step_tensor_cache.get(id(group)) |
| 16 | - amsgrad: bool, | 15 | + if step_tensor is None or step_tensor.device.index != device: |
| 17 | - beta1: float, | 16 | + step_tensor = torch.empty((), dtype=torch.int64, device=device) |
| 18 | - beta2: float, | 17 | + step_tensor_cache[id(group)] = step_tensor |
| 19 | - lr: float, | 18 | + step_tensor.fill_(group['step']) |
| 20 | - weight_decay: float, | 19 | + return step_tensor |
🔵 Low Priority 上一轮审查已指出
本次 diff 在 建议:将 ![]() ![]() 不准确? | |||
| 21 | - eps: float, | 20 | + |
| 22 | - maximize: bool): | 21 | + |
| 22 | +def adamw( | ||
| 23 | + params: List[Tensor], | ||
| 24 | + grads: List[Tensor], | ||
| 25 | + exp_avgs: List[Tensor], | ||
| 26 | + exp_avg_sqs: List[Tensor], | ||
| 27 | + max_exp_avg_sqs: List[Tensor], | ||
| 28 | + step_tensor: Tensor, | ||
| 29 | + *, | ||
| 30 | + amsgrad: bool, | ||
| 31 | + beta1: float, | ||
| 32 | + beta2: float, | ||
| 33 | + lr: float, | ||
| 34 | + weight_decay: float, | ||
| 35 | + eps: float, | ||
| 36 | + maximize: bool, | ||
| 37 | +): | ||
| 23 | r"""Functional API that performs AdamW algorithm computation. | 38 | r"""Functional API that performs AdamW algorithm computation. |
| 24 | See :class:`~torch.optim.AdamW` for details. | 39 | See :class:`~torch.optim.AdamW` for details. |
| 25 | """ | 40 | """ |
| @@ -42,33 +57,39 @@ def adamw(params: List[Tensor], | |||
| 42 | beta2=beta2, | 57 | beta2=beta2, |
| 43 | weight_decay=weight_decay, | 58 | weight_decay=weight_decay, |
| 44 | eps=eps, | 59 | eps=eps, |
| 45 | - maximize=maximize | 60 | + maximize=maximize, |
| 46 | ) | 61 | ) |
| 47 | 62 | ||
| 48 | 63 | ||
| 49 | class AdamW(Optimizer): | 64 | class AdamW(Optimizer): |
| 50 | - def __init__(self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, | 65 | + def __init__( |
| 51 | - weight_decay=1e-2, amsgrad=False, *, maximize: bool = False): | 66 | + self, params, lr=1e-3, betas=(0.9, 0.999), eps=1e-8, weight_decay=1e-2, amsgrad=False, *, maximize: bool = False |
| 52 | - if not 0.0 <= lr: | 67 | + ): |
| 68 | + if 0.0 > lr: | ||
| 53 | raise ValueError("Invalid learning rate: {}".format(lr)) | 69 | raise ValueError("Invalid learning rate: {}".format(lr)) |
| 54 | - if not 0.0 <= eps: | 70 | + if 0.0 > eps: |
| 55 | raise ValueError("Invalid epsilon value: {}".format(eps)) | 71 | raise ValueError("Invalid epsilon value: {}".format(eps)) |
| 56 | if not 0.0 <= betas[0] < 1.0: | 72 | if not 0.0 <= betas[0] < 1.0: |
| 57 | raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) | 73 | raise ValueError("Invalid beta parameter at index 0: {}".format(betas[0])) |
| 58 | if not 0.0 <= betas[1] < 1.0: | 74 | if not 0.0 <= betas[1] < 1.0: |
| 59 | raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) | 75 | raise ValueError("Invalid beta parameter at index 1: {}".format(betas[1])) |
| 60 | - if not 0.0 <= weight_decay: | 76 | + if 0.0 > weight_decay: |
| 61 | raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) | 77 | raise ValueError("Invalid weight_decay value: {}".format(weight_decay)) |
| 62 | - defaults = dict(lr=lr, betas=betas, eps=eps, | 78 | + defaults = dict(lr=lr, betas=betas, eps=eps, weight_decay=weight_decay, amsgrad=amsgrad, maximize=maximize) |
| 63 | - weight_decay=weight_decay, amsgrad=amsgrad, maximize=maximize) | 79 | + super().__init__(params, defaults) |
| 64 | - super(AdamW, self).__init__(params, defaults) | 80 | + self._step_tensor_cache = {} |
| 65 | 81 | ||
| 66 | def __setstate__(self, state): | 82 | def __setstate__(self, state): |
| 67 | - super(AdamW, self).__setstate__(state) | 83 | + super().__setstate__(state) |
| 68 | for group in self.param_groups: | 84 | for group in self.param_groups: |
| 69 | group.setdefault('amsgrad', False) | 85 | group.setdefault('amsgrad', False) |
| 70 | group.setdefault('maximize', False) | 86 | group.setdefault('maximize', False) |
| 71 | 87 | ||
| 88 | + def load_state_dict(self, state_dict): | ||
| 89 | + result = super().load_state_dict(state_dict) | ||
| 90 | + self._step_tensor_cache = {} | ||
| 91 | + return result | ||
| 92 | + | ||
| 72 | 93 | ||
| 73 | def step(self, closure=None): | 94 | def step(self, closure=None): |
| 74 | loss = None | 95 | loss = None |
| @@ -81,18 +102,15 @@ class AdamW(Optimizer): | |||
| 81 | grads = [] | 102 | grads = [] |
| 82 | exp_avgs = [] | 103 | exp_avgs = [] |
| 83 | exp_avg_sqs = [] | 104 | exp_avg_sqs = [] |
| 84 | - state_sums = [] | ||
| 85 | max_exp_avg_sqs = [] | 105 | max_exp_avg_sqs = [] |
| 86 | - state_steps = [] | ||
| 87 | amsgrad = group['amsgrad'] | 106 | amsgrad = group['amsgrad'] |
| 88 | beta1, beta2 = group['betas'] | 107 | beta1, beta2 = group['betas'] |
| 89 | 108 | ||
| 90 | if 'step' in group: | 109 | if 'step' in group: |
| 91 | - group['step'] += 1 | 110 | + group['step'] = _to_step_int(group['step']) + 1 |
| 92 | - if group['step'].is_cpu: | ||
| 93 | - group['step'] = group['step'].cuda() | ||
| 94 | else: | 111 | else: |
| 95 | - group['step'] = torch.tensor(1, dtype=torch.int64, device=torch.cuda.current_device()) | 112 | + group['step'] = 1 |
| 113 | + step_tensor = _get_step_tensor(self, group) | ||
| 96 | 114 | ||
| 97 | for p in group['params']: | 115 | for p in group['params']: |
| 98 | if p.grad is None: | 116 | if p.grad is None: |
| @@ -120,18 +138,20 @@ class AdamW(Optimizer): | |||
| 120 | if amsgrad: | 138 | if amsgrad: |
| 121 | max_exp_avg_sqs.append(state['max_exp_avg_sq']) | 139 | max_exp_avg_sqs.append(state['max_exp_avg_sq']) |
| 122 | 140 | ||
| 123 | - adamw(params_with_grad, | 141 | + adamw( |
| 124 | - grads, | 142 | + params_with_grad, |
| 125 | - exp_avgs, | 143 | + grads, |
| 126 | - exp_avg_sqs, | 144 | + exp_avgs, |
| 127 | - max_exp_avg_sqs, | 145 | + exp_avg_sqs, |
| 128 | - group['step'], | 146 | + max_exp_avg_sqs, |
| 129 | - amsgrad=amsgrad, | 147 | + step_tensor, |
| 130 | - beta1=beta1, | 148 | + amsgrad=amsgrad, |
| 131 | - beta2=beta2, | 149 | + beta1=beta1, |
| 132 | - lr=group['lr'], | 150 | + beta2=beta2, |
| 133 | - weight_decay=group['weight_decay'], | 151 | + lr=group['lr'], |
| 134 | - eps=group['eps'], | 152 | + weight_decay=group['weight_decay'], |
| 135 | - maximize=group['maximize']) | 153 | + eps=group['eps'], |
| 154 | + maximize=group['maximize'], | ||
| 155 | + ) | ||
| 136 | 156 | ||
| 137 | - return loss | 157 | + return loss |
| @@ -115,6 +115,31 @@ class MindSpeedTELayerNormColumnParallelLinear(torch.nn.Module): | |||
| 115 | if self.allreduce_dgrad and self.sequence_parallel: | 115 | if self.allreduce_dgrad and self.sequence_parallel: |
| 116 | raise RuntimeError("`allreduce_dgrad` and `sequence_parallel` cannot be enabled at the same time.") | 116 | raise RuntimeError("`allreduce_dgrad` and `sequence_parallel` cannot be enabled at the same time.") |
| 117 | 117 | ||
| 118 | + if self.sequence_parallel and tp_size <= 1: | ||
| 119 | + warnings.warn( | ||
| 120 | + "`sequence_parallel` is set to `True`, but tensor model parallel size " | ||
| 121 | + f"is {tp_size}. Disabling sequence parallel." | ||
| 122 | + ) | ||
| 123 | + self.sequence_parallel = False | ||
| 124 | + | ||
| 125 | + # Norm init spec. | ||
| 126 | + if self.config.normalization not in ['LayerNorm', 'RMSNorm']: | ||
| 127 | + raise AssertionError('Unsupported normalization type {}!'.format(self.config.normalization)) | ||
| 128 | + | ||
| 129 | + layer_norm_weight = torch.nn.Parameter( | ||
| 130 | + torch.ones(self.input_size, device='npu', dtype=self.config.params_dtype) | ||
| 131 | + ) | ||
| 132 | + self.register_parameter("layer_norm_weight", layer_norm_weight) | ||
| 133 | + setattr(self.layer_norm_weight, 'sequence_parallel', self.sequence_parallel) | ||
| 134 | + | ||
| 135 | + self.register_parameter("layer_norm_bias", None) | ||
| 136 | + if self.config.normalization != 'RMSNorm': | ||
| 137 | + layer_norm_bias = torch.nn.Parameter( | ||
| 138 | + torch.zeros(self.input_size, device='npu', dtype=self.config.params_dtype) | ||
| 139 | + ) | ||
| 140 | + setattr(layer_norm_bias, 'sequence_parallel', self.sequence_parallel) | ||
| 141 | + self.layer_norm_bias = layer_norm_bias | ||
| 142 | + | ||
| 118 | # Because skip_weight_param_allocation is not supported in TE, always do weight initialize. | 143 | # Because skip_weight_param_allocation is not supported in TE, always do weight initialize. |
| 119 | if config.use_cpu_initialization: | 144 | if config.use_cpu_initialization: |
| 120 | self.weight = torch.nn.Parameter( | 145 | self.weight = torch.nn.Parameter( |
| @@ -173,33 +198,9 @@ class MindSpeedTELayerNormColumnParallelLinear(torch.nn.Module): | |||
| 173 | else: | 198 | else: |
| 174 | self.register_parameter('bias', None) | 199 | self.register_parameter('bias', None) |
| 175 | 200 | ||
| 176 | - if self.sequence_parallel and tp_size <= 1: | ||
| 177 | - warnings.warn( | ||
| 178 | - "`sequence_parallel` is set to `True`, but tensor model parallel size " | ||
| 179 | - f"is {tp_size}. Disabling sequence parallel." | ||
| 180 | - ) | ||
| 181 | - self.sequence_parallel = False | ||
| 182 | - | ||
| 183 | # Forward impl settings without ascend-mc2. | 201 | # Forward impl settings without ascend-mc2. |
| 184 | self._linear_forward_impl = linear_with_grad_accumulation_and_async_allreduce | 202 | self._linear_forward_impl = linear_with_grad_accumulation_and_async_allreduce |
| 185 | 203 | ||
| 186 | - # Norm init spec. | ||
| 187 | - if self.config.normalization not in ['LayerNorm', 'RMSNorm']: | ||
| 188 | - raise AssertionError('Unsupported normalization type {}!'.format(self.config.normalization)) | ||
| 189 | - | ||
| 190 | - layer_norm_weight = torch.nn.Parameter( | ||
| 191 | - torch.ones(self.input_size, device='npu', dtype=self.config.params_dtype) | ||
| 192 | - ) | ||
| 193 | - self.register_parameter("layer_norm_weight", layer_norm_weight) | ||
| 194 | - setattr(self.layer_norm_weight, 'sequence_parallel', self.sequence_parallel) | ||
| 195 | - | ||
| 196 | - self.register_parameter("layer_norm_bias", None) | ||
| 197 | - if self.config.normalization != 'RMSNorm': | ||
| 198 | - layer_norm_bias = torch.nn.Parameter( | ||
| 199 | - torch.zeros(self.input_size, device='npu', dtype=self.config.params_dtype) | ||
| 200 | - ) | ||
| 201 | - setattr(layer_norm_bias, 'sequence_parallel', self.sequence_parallel) | ||
| 202 | - self.layer_norm_bias = layer_norm_bias | ||
| 203 | self.te_return_bias = self.skip_bias_add and bias | 204 | self.te_return_bias = self.skip_bias_add and bias |
| 204 | 205 | ||
| 205 | def _layernorm(self, inp): | 206 | def _layernorm(self, inp): |


🔵 Low Priority
_get_step_tensor函数(mindspeed/core/optimizer/adamw.py第 15-26 行,mindspeed/optimizer/adamw.py同样位置)使用id(group)作为_step_tensor_cache字典的键,缓存标量 step tensor 以避免每次 step 都创建新 tensor。但此缓存没有清理机制:如果 param_groups 在运行时被替换(如加载 state_dict 后),旧 group 对应的 tensor 将永久保留在缓存中。虽然标量 tensor 内存开销很小(int64 × 1 = 8 bytes),但长期运行(如训练数天)且频繁替换 group 的场景下可能累积。此外,若旧 group 的
id()被 Python 回收后重用给新 group,会错误命中旧缓存(概率极低但因 Python 的对象 id 复用而存在理论可能)。建议:可选改进:在
AdamW.load_state_dict()被调用后清空_step_tensor_cache(因为 param_groups 可能已被替换)。或者使用weakref引用 group 对象来避免缓存泄漏。由于该问题触发条件苛刻且影响很小,此项为可选建议。