已合并
main_grad not found bugfix for lora #2539
mhh001创建于 2025年4月16日
main_grad not found bugfix for lora #2539
已合并
从refs/pull/2539/head合入到2.0.0
共 2 个文件变更+55-2
| @@ -4,15 +4,64 @@ | |||
| 4 | from typing import List, Optional | 4 | from typing import List, Optional |
| 5 | 5 | ||
| 6 | import torch | 6 | import torch |
| 7 | +from torch._utils import _flatten_dense_tensors, _unflatten_dense_tensors | ||
| 7 | 8 | ||
| 8 | from megatron.core import parallel_state | 9 | from megatron.core import parallel_state |
| 9 | from megatron.core.distributed.finalize_model_grads import _allreduce_layernorm_grads, _allreduce_embedding_grads | 10 | from megatron.core.distributed.finalize_model_grads import _allreduce_layernorm_grads, _allreduce_embedding_grads |
| 10 | from megatron.core.transformer.transformer_config import TransformerConfig | 11 | from megatron.core.transformer.transformer_config import TransformerConfig |
| 11 | from megatron.core.utils import get_attr_wrapped_model, get_model_config | 12 | from megatron.core.utils import get_attr_wrapped_model, get_model_config |
| 12 | from megatron.training import get_args | 13 | from megatron.training import get_args |
| 14 | +from mindspeed.core.tensor_parallel.comm_group_api import TPXCollectiveComm | ||
| 13 | from mindspeed_llm.core.transformer.moe.moe_utils import get_updated_expert_bias | 15 | from 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 | + | ||
| 16 | def _allreduce_word_embedding_grads(model: List[torch.nn.Module], config: TransformerConfig): | 65 | def _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_grad | 93 | 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_grad | 100 | 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.core | 180 | import megatron.core |
| 181 | megatron.core.jit.jit_fuser = dummy_jit | 181 | megatron.core.jit.jit_fuser = dummy_jit |
| 182 | - from mindspeed.core.tensor_parallel.tp_2d.norm_factory import _allreduce_layernorm_grads_wrapper | 182 | + 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 embedding | 185 | # Mtp share embedding |
| 186 | from mindspeed_llm.core.distributed.finalize_model_grads import _allreduce_word_embedding_grads | 186 | 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', |