已合并
main_grad not found bugfix for lora #2539
mhh001创建于 2025年4月16日
main_grad not found bugfix for lora #2539
已合并
mhh001创建于 2025年4月16日
refs/pull/2539/head合入到2.0.0
2 个文件变更+55-2
@@ -4,15 +4,64 @@
4from typing import List, Optional4from typing import List, Optional
5 5 
6import torch6import torch
7+from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors
7 8 
8from megatron.core import parallel_state9from megatron.core import parallel_state
9from megatron.core.distributed.finalize_model_grads import _allreduce_layernorm_grads, _allreduce_embedding_grads10from megatron.core.distributed.finalize_model_grads import _allreduce_layernorm_grads, _allreduce_embedding_grads
10from megatron.core.transformer.transformer_config import TransformerConfig11from megatron.core.transformer.transformer_config import TransformerConfig
11from megatron.core.utils import get_attr_wrapped_model, get_model_config12from megatron.core.utils import get_attr_wrapped_model, get_model_config
12from megatron.training import get_args13from megatron.training import get_args
14+from mindspeed.core.tensor_parallel.comm_group_api import TPXCollectiveComm
13from mindspeed_llm.core.transformer.moe.moe_utils import get_updated_expert_bias15from mindspeed_llm.core.transformer.moe.moe_utils import get_updated_expert_bias
14 16 
15 17 
18+def allreduce_layernorm_grads(model: List[torch.nn.Module], config: TransformerConfig):
19+ """
20+ All-reduce layernorm grads (for sequence parallelism).
21+ """
22+ 
23+ # All-reduce layernorm parameters across model parallel nodes
24+ # when sequence parallelism is used
25+ if parallel_state.get_tensor_model_parallel_world_size() > 1 and (
26+ config.sequence_parallel or config.qk_layernorm
27+ ):
28+ grads = []
29+ for model_chunk in model:
30+ for name, param in get_attr_wrapped_model(model_chunk, 'named_parameters')():
31+ if not param.requires_grad:
32+ continue
33+ if (
34+ param.requires_grad
35+ and getattr(param, 'sequence_parallel', False)
36+ or 'q_layernorm' in name
37+ or 'k_layernorm' in name
38+ ):
39+ grad = param.main_grad
40+ grads.append(grad.data)
41+ if grads:
42+ coalesced = _flatten_dense_tensors(grads)
43+ torch.distributed.all_reduce(
44+ coalesced, group=parallel_state.get_tensor_model_parallel_group()
45+ )
46+ for buf, synced in zip(grads, _unflatten_dense_tensors(coalesced, grads)):
47+ buf.copy_(synced)
48+ 
49+ layer_norm_2d_grads = []
50+ for model_chunk in model:
51+ for name, param in get_attr_wrapped_model(model_chunk, "named_parameters")():
52+ if param.requires_grad and getattr(param, "2d_tp", False):
53+ layer_norm_2d_grad = param.main_grad
54+ layer_norm_2d_grads.append(layer_norm_2d_grad.data)
55+ 
56+ if layer_norm_2d_grads:
57+ coalesced = _flatten_dense_tensors(layer_norm_2d_grads)
58+ torch.distributed.all_reduce(coalesced, group=TPXCollectiveComm.get_comm_group())
59+ for buf, synced in zip(
60+ layer_norm_2d_grads, _unflatten_dense_tensors(coalesced, layer_norm_2d_grads)
61+ ):
62+ buf.copy_(synced)
63+ 
64+ 
16def _allreduce_word_embedding_grads(model: List[torch.nn.Module], config: TransformerConfig):65def _allreduce_word_embedding_grads(model: List[torch.nn.Module], config: TransformerConfig):
17 """66 """
18 All-reduce word embedding grads.67 All-reduce word embedding grads.
@@ -39,11 +88,15 @@ def _allreduce_word_embedding_grads(model: List[torch.nn.Module], config: Transf
39 model_module = get_attr_wrapped_model(model_module, 'pre_process', return_model_obj=True)88 model_module = get_attr_wrapped_model(model_module, 'pre_process', return_model_obj=True)
40 if model_module.share_embeddings_and_output_weights:89 if model_module.share_embeddings_and_output_weights:
41 weight = model_module.shared_embedding_or_output_weight()90 weight = model_module.shared_embedding_or_output_weight()
91+ if not weight.requires_grad:
92+ return
42 grad = weight.main_grad93 grad = weight.main_grad
43 torch.distributed.all_reduce(grad, group=parallel_state.get_embedding_group())94 torch.distributed.all_reduce(grad, group=parallel_state.get_embedding_group())
44 if hasattr(model_module,95 if hasattr(model_module,
45 "share_mtp_embedding_and_output_weight") and model_module.share_mtp_embedding_and_output_weight:96 "share_mtp_embedding_and_output_weight") and model_module.share_mtp_embedding_and_output_weight:
46 weight = model_module.shared_embedding_weight()97 weight = model_module.shared_embedding_weight()
98+ if not weight.requires_grad:
99+ return
47 grad = weight.main_grad100 grad = weight.main_grad
48 torch.distributed.all_reduce(grad, group=parallel_state.get_embedding_group())101 torch.distributed.all_reduce(grad, group=parallel_state.get_embedding_group())
49 102 
@@ -179,9 +179,9 @@ class CoreAdaptation(MegatronAdaptationABC):
179 def patch_core_distributed(self):179 def patch_core_distributed(self):
180 import megatron.core180 import megatron.core
181 megatron.core.jit.jit_fuser = dummy_jit181 megatron.core.jit.jit_fuser = dummy_jit
182- from mindspeed.core.tensor_parallel.tp_2d.norm_factory import _allreduce_layernorm_grads_wrapper182+ from mindspeed_llm.core.distributed.finalize_model_grads import allreduce_layernorm_grads
183 MegatronAdaptation.register('megatron.core.distributed.finalize_model_grads._allreduce_layernorm_grads',183 MegatronAdaptation.register('megatron.core.distributed.finalize_model_grads._allreduce_layernorm_grads',
184- _allreduce_layernorm_grads_wrapper)184+ allreduce_layernorm_grads)
185 # Mtp share embedding185 # Mtp share embedding
186 from mindspeed_llm.core.distributed.finalize_model_grads import _allreduce_word_embedding_grads186 from mindspeed_llm.core.distributed.finalize_model_grads import _allreduce_word_embedding_grads
187 MegatronAdaptation.register('megatron.core.distributed.finalize_model_grads._allreduce_word_embedding_grads',187 MegatronAdaptation.register('megatron.core.distributed.finalize_model_grads._allreduce_word_embedding_grads',