已合并
fix: TE LayerNormLinear init weight order align NVTE #3569
fix: TE LayerNormLinear init weight order align NVTE #3569
已合并
clc2025创建于 6月23日
8 个文件变更+322-181
@@ -9,6 +9,7 @@ from torch import _C
9from torch_npu.npu import _lazy_call, device as device_ctx_manager9from torch_npu.npu import _lazy_call, device as device_ctx_manager
10from megatron.core.optimizer.cpu_offloading import HybridDeviceOptimizer10from megatron.core.optimizer.cpu_offloading import HybridDeviceOptimizer
11from megatron.core.optimizer.distrib_optimizer import HAVE_APEX_OR_TE11from megatron.core.optimizer.distrib_optimizer import HAVE_APEX_OR_TE
12+from mindspeed.core.optimizer.utils import _to_step_int
12from mindspeed.core.tensor_parallel.tp_2d.group_api_2d import TPYCollectiveComm13from mindspeed.core.tensor_parallel.tp_2d.group_api_2d import TPYCollectiveComm
13from mindspeed.core.tensor_parallel.tp_2d.layernorm_2d import LayerNorm2D14from mindspeed.core.tensor_parallel.tp_2d.layernorm_2d import LayerNorm2D
14from mindspeed.core.tensor_parallel.tp_2d.rms_norm_2d import RMSNorm2D15from 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_helpers57 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 
62def add_layer_norm_sp_support(config, instance):66def 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- 
73class PTNorm:76class 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 implementation87 # using apex implementation
86 from megatron.core.fusions.fused_layer_norm import FusedLayerNorm88 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 implementation92 # using torch implementation
@@ -99,10 +102,11 @@ class PTNorm:
99 instance.use_fused_rmsnorm = False102 instance.use_fused_rmsnorm = False
100 else:103 else:
101 from mindspeed.core.fusions.fused_rms_norm import RMSNorm104 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 = True107 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 instance111 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 device126 return device
127+ 
123 return wrapper128 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) = bucket148 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_data150+ (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_groups224 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 docstring250 # 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 AdamW287 # 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 @@
1from typing import List, Optional, Tuple, Union1from typing import List, Optional, Tuple, Union
2import torch2import torch
3-import torch_npu
4from torch import Tensor3from torch import Tensor
5from torch.optim.optimizer import Optimizer4from torch.optim.optimizer import Optimizer
6from torch.optim.adamw import AdamW as TorchAdamW5from 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
atomgit-bot
atomgit-botatomgit-bot6月23日

🔵 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 对象来避免缓存泄漏。由于该问题触发条件苛刻且影响很小,此项为可选建议。

likedislike
不准确?
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=maximize61+ 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 
78class AdamW(Optimizer):96class 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 @torch.no_grad()125 @torch.no_grad()
102 def step(self, closure=None):126 def step(self, closure=None):
103 loss = None127 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'] += 1142+ 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 loss189 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
atomgit-bot
atomgit-botatomgit-bot6月23日

🟡 Medium Priority

mindspeed/core/optimizer/optimizer.py 第 29 行,state_dict = states[idx] if states else None 使用 enumerate(self.chained_optimizers) 的原始索引 idx 访问 states 列表。

当某个 optimizer 不满足 hasattr(optimizer, 'load_parameter_state_from_dp_zero') 条件被 continue 跳过时(第 23-24 行),idx 仍然递增(因为是 enumerate 的原始索引),而 states 列表的格式可能不包含被跳过的 optimizer 对应的条目。这会导致后续 optimizer 取到错误的 state_dict(取到相邻 optimizer 的状态),或在 states 列表长度不足时引发 IndexError

此问题在上一轮审查中已提出,当前 diff 中未修复。虽然在实际使用中所有 chained_optimizers 通常都是 DistributedOptimizer 实例(均有 load_parameter_state_from_dp_zero 方法),但代码逻辑上存在此风险。

建议:建议引入一个独立的计数器 state_idx 来跟踪实际匹配的 optimizer 数量,而非使用 enumerate 的原始索引。或确认保存侧 states 文件的格式始终为所有 chained_optimizers(包括无 load_parameter_state_from_dp_zero 的 optimizer)保存对应条目。如果格式确实包含所有条目,可以添加注释说明。

likedislike
不准确?
clc2025
6月24日 评论:
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 = None20 args.virtual_pipeline_model_parallel_size = None
18 args.overlap_p2p_comm = False21 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.training40+ 
36- only_mcore = False41+ 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 patches48 # 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 patches71 # initialization patches
55 from mindspeed.core.megatron_basic.megatron_basic import _set_cuda_rng_state72 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 patches76 # norm patches
59 from mindspeed.core.megatron_basic.megatron_basic import PTNorm77 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_sync85 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_wrapper86 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_wrapper87 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 params99 # Currently, it is not supported to Cast shard fp32 main params to fp8 model params
74 from mindspeed.core.fp8_utils import quantize_param_shard100 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 step105 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 save109 # avoid async save
82 from mindspeed.core.megatron_basic.megatron_basic import preload_tensors110 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 task117 # avoid incorrect weight_decay override in resume task
86 from mindspeed.core.megatron_basic.megatron_basic import dist_optim_load_state_dict118 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 patch137 # args parser patch
91- from mindspeed.core.megatron_basic.arguments_basic import parse_args_wrapper, validate_args_wrapper, print_args_wrapper138+ 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 patches150 # initialization patches
99 from mindspeed.core.megatron_basic.megatron_basic import _compile_dependencies, get_device_wrapper151 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_version156 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, Union1+from typing import List
2import torch2import torch
3-import torch_npu
4from torch import Tensor3from torch import Tensor
5from torch.optim.optimizer import Optimizer4from torch.optim.optimizer import Optimizer
6-from torch.optim.adamw import AdamW as TorchAdamW5+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
atomgit-bot
atomgit-botatomgit-bot6月23日

🔵 Low Priority

上一轮审查已指出 _to_step_int_get_step_tensor 在两个文件中完全重复定义:

  • mindspeed/optimizer/adamw.py 第 9-12 行 / 第 15-26 行(本次 diff 新增)
  • mindspeed/core/optimizer/adamw.py 第 9-12 行 / 第 15-26 行(已存在)

本次 diff 在 mindspeed/optimizer/adamw.py 中添加了这两个函数,与 mindspeed/core/optimizer/adamw.py 中的实现完全相同。两处定义独立维护,未来任一处的修改可能遗漏同步另一处,导致行为不一致。建议将这两个函数提取到共享工具模块(如 mindspeed/core/megatron_basic/megatron_basic.py)中,让两边从同一来源导入。

建议:将 _to_step_int_get_step_tensor 提取到一个共享工具模块中,让 mindspeed/optimizer/adamw.pymindspeed/core/optimizer/adamw.py 都从同一处导入。

likedislike
不准确?
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=maximize60+ maximize=maximize,
46 )61 )
47 62 
48 63 
49class AdamW(Optimizer):64class 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 @torch.no_grad()93 @torch.no_grad()
73 def step(self, closure=None):94 def step(self, closure=None):
74 loss = None95 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'] += 1110+ 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 loss157+ 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_allreduce202 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 bias204 self.te_return_bias = self.skip_bias_add and bias
204 205 
205 def _layernorm(self, inp):206 def _layernorm(self, inp):