已合并
[mindspore][patch][master]del patch #3696
王一博创建于 2025年11月15日
[mindspore][patch][master]del patch #3696
已合并
共 2 个文件变更+2-197
| @@ -82,13 +82,6 @@ def _patch_model_and_embedding(): | |||
| 82 | 'mindspeed_llm.tasks.common.yarn_rope.YarnRotaryPositionEmbedding.yarn_linear_ramp_mask', | 82 | 'mindspeed_llm.tasks.common.yarn_rope.YarnRotaryPositionEmbedding.yarn_linear_ramp_mask', |
| 83 | yarn_linear_ramp_mask) | 83 | yarn_linear_ramp_mask) |
| 84 | 84 | ||
| 85 | - from mindspeed.mindspore.core.transformer.module import fp32_to_float16 | ||
| 86 | - MindSporeAdaptation.register_patch('megatron.core.transformer.module.fp32_to_float16', fp32_to_float16) | ||
| 87 | - | ||
| 88 | - from mindspeed.mindspore.legacy.model.module import fp32_to_float16, float16_to_fp32 | ||
| 89 | - MindSporeAdaptation.register_patch('megatron.legacy.model.module.fp32_to_float16', fp32_to_float16) | ||
| 90 | - MindSporeAdaptation.register_patch('megatron.legacy.model.module.float16_to_fp32', float16_to_fp32) | ||
| 91 | - | ||
| 92 | from mindspeed_llm.mindspore.core.models.common.embeddings.rotary_pos_embedding import apply_llama3_scaling | 85 | from mindspeed_llm.mindspore.core.models.common.embeddings.rotary_pos_embedding import apply_llama3_scaling |
| 93 | MindSporeAdaptation.register_patch( | 86 | MindSporeAdaptation.register_patch( |
| 94 | 'mindspeed_llm.core.models.common.embeddings.rotary_pos_embedding.apply_llama3_scaling', apply_llama3_scaling) | 87 | 'mindspeed_llm.core.models.common.embeddings.rotary_pos_embedding.apply_llama3_scaling', apply_llama3_scaling) |
| @@ -156,10 +149,10 @@ def _patch_optimizer_and_training(args): | |||
| 156 | # share memory | 149 | # share memory |
| 157 | if args.enable_share_memory: | 150 | if args.enable_share_memory: |
| 158 | from ..mindspore.tasks.dataset.shared_memory_manager import SharedMemoryManager | 151 | from ..mindspore.tasks.dataset.shared_memory_manager import SharedMemoryManager |
| 159 | - MindSporeAdaptation.register( | 152 | + MindSporeAdaptation.register_patch( |
| 160 | 'mindspeed_llm.tasks.dataset.shared_memory_manager.SharedMemoryManager', SharedMemoryManager) | 153 | 'mindspeed_llm.tasks.dataset.shared_memory_manager.SharedMemoryManager', SharedMemoryManager) |
| 161 | from ..mindspore.training.utils import _compute_actual_seq_len | 154 | from ..mindspore.training.utils import _compute_actual_seq_len |
| 162 | - MindSporeAdaptation.register( | 155 | + MindSporeAdaptation.register_patch( |
| 163 | 'mindspeed_llm.training.utils._compute_actual_seq_len', _compute_actual_seq_len) | 156 | 'mindspeed_llm.training.utils._compute_actual_seq_len', _compute_actual_seq_len) |
| 164 | 157 | ||
| 165 | # Optimizer: load and save parameter | 158 | # Optimizer: load and save parameter |
| @@ -257,7 +250,6 @@ def _patch_fused_operators(args): | |||
| 257 | # Matmul add ops | 250 | # Matmul add ops |
| 258 | from mindspeed.mindspore.ops.npu_matmul_add import npu_matmul_add_fp32 | 251 | from mindspeed.mindspore.ops.npu_matmul_add import npu_matmul_add_fp32 |
| 259 | MindSporeAdaptation.register_patch('fused_weight_gradient_mlp_cuda.wgrad_gemm_accum_fp32', npu_matmul_add_fp32) | 252 | MindSporeAdaptation.register_patch('fused_weight_gradient_mlp_cuda.wgrad_gemm_accum_fp32', npu_matmul_add_fp32) |
| 260 | - MindSporeAdaptation.register_patch('mindspeed.ops.npu_matmul_add.npu_matmul_add_fp32', npu_matmul_add_fp32) | ||
| 261 | 253 | ||
| 262 | # Fused AdamW v2 | 254 | # Fused AdamW v2 |
| 263 | from torch import npu_apply_fused_adamw_v2 | 255 | from torch import npu_apply_fused_adamw_v2 |
| @@ -2,20 +2,7 @@ | |||
| 2 | # Copyright (c) 2025, Huawei Technologies Co., Ltd. All rights reserved. | 2 | # Copyright (c) 2025, Huawei Technologies Co., Ltd. All rights reserved. |
| 3 | 3 | ||
| 4 | """General utilities.""" | 4 | """General utilities.""" |
| 5 | -import logging | ||
| 6 | -from itertools import takewhile | ||
| 7 | -import torch | ||
| 8 | import numpy as np | 5 | import numpy as np |
| 9 | -from megatron.training import get_args | ||
| 10 | -from megatron.core import mpu | ||
| 11 | -import acl | ||
| 12 | -from mindspeed_llm.training.utils import (get_sharedmem_mgr, BASE_SHM_NAME, compute_actual_seq_len, | ||
| 13 | - set_mtp_position_ids, regenerate_position_ids) | ||
| 14 | - | ||
| 15 | -try: | ||
| 16 | - from mindspeed.core.pipeline_parallel.dualpipev.dualpipev_schedules import get_post_process_flag | ||
| 17 | -except ImportError as e: | ||
| 18 | - logging.warning(f"Import failed: {e}") | ||
| 19 | 6 | ||
| 20 | 7 | ||
| 21 | def _compute_actual_seq_len(origin_seq): | 8 | def _compute_actual_seq_len(origin_seq): |
| @@ -28,177 +15,3 @@ def _compute_actual_seq_len(origin_seq): | |||
| 28 | 15 | ||
| 29 | res.append(len(seq)) | 16 | res.append(len(seq)) |
| 30 | return res | 17 | return res |
| 31 | - | ||
| 32 | - | ||
| 33 | -def get_batch_on_this_tp_rank(data_iterator): | ||
| 34 | - args = get_args() | ||
| 35 | - | ||
| 36 | - def _broadcast(item): | ||
| 37 | - if item is not None: | ||
| 38 | - torch.distributed.broadcast(item, mpu.get_tensor_model_parallel_src_rank(), | ||
| 39 | - group=mpu.get_tensor_model_parallel_group()) | ||
| 40 | - | ||
| 41 | - shm_manager = None | ||
| 42 | - actual_seq_len = None | ||
| 43 | - if args.enable_share_memory: | ||
| 44 | - shm_manager = get_sharedmem_mgr(BASE_SHM_NAME, args.micro_batch_size * args.seq_length) | ||
| 45 | - | ||
| 46 | - if mpu.get_tensor_model_parallel_rank() == 0: | ||
| 47 | - if data_iterator is not None: | ||
| 48 | - data = next(data_iterator) | ||
| 49 | - else: | ||
| 50 | - data = None | ||
| 51 | - | ||
| 52 | - if args.enable_share_memory and shm_manager is not None: | ||
| 53 | - position_ids = data["position_ids"] | ||
| 54 | - actual_seq_len = compute_actual_seq_len(position_ids) | ||
| 55 | - shm_manager.write(actual_seq_len) | ||
| 56 | - | ||
| 57 | - if '910B' not in acl.get_soc_name() and args.mtp_num_layers and get_post_process_flag(): | ||
| 58 | - from mindspeed_llm.core.transformer.multi_token_prediction import roll_tensor | ||
| 59 | - position_ids_mtp = [] | ||
| 60 | - cur_position_id = data["position_ids"] | ||
| 61 | - for _ in range(args.mtp_num_layers): | ||
| 62 | - cur_position_id, _ = roll_tensor(cur_position_id, shifts=-1, dims=-1) | ||
| 63 | - cur_position_id = regenerate_position_ids(cur_position_id, 1) | ||
| 64 | - position_ids_mtp.append(cur_position_id) | ||
| 65 | - set_mtp_position_ids((position_ids_mtp, shm_manager)) | ||
| 66 | - | ||
| 67 | - if args.return_document_ids and mpu.get_context_parallel_rank() == 0 and mpu.get_pipeline_model_parallel_rank() == 0: | ||
| 68 | - document_ids = [ | ||
| 69 | - [x.item() for x in takewhile(lambda y: y.item() != -100, row)] | ||
| 70 | - for row in data['document_ids'] | ||
| 71 | - ] | ||
| 72 | - data_idx = [ | ||
| 73 | - [x.item() for x in takewhile(lambda y: y.item() != -100, row)] | ||
| 74 | - for row in data['idx'] | ||
| 75 | - ] | ||
| 76 | - | ||
| 77 | - data.pop("document_ids", None) | ||
| 78 | - data.pop("idx", None) | ||
| 79 | - | ||
| 80 | - batch = { | ||
| 81 | - 'tokens': data["tokens"], | ||
| 82 | - 'labels': data["labels"], | ||
| 83 | - 'loss_mask': data["loss_mask"], | ||
| 84 | - 'attention_mask': None if "attention_mask" not in data else data["attention_mask"], | ||
| 85 | - 'position_ids': data["position_ids"], | ||
| 86 | - 'document_ids': document_ids, | ||
| 87 | - 'idx': data_idx | ||
| 88 | - } | ||
| 89 | - else: | ||
| 90 | - batch = { | ||
| 91 | - 'tokens': data["tokens"], | ||
| 92 | - 'labels': data["labels"], | ||
| 93 | - 'loss_mask': data["loss_mask"], | ||
| 94 | - 'attention_mask': None if "attention_mask" not in data else data["attention_mask"], | ||
| 95 | - 'position_ids': data["position_ids"] | ||
| 96 | - } | ||
| 97 | - if args.pipeline_model_parallel_size == 1: | ||
| 98 | - _broadcast(batch['tokens']) | ||
| 99 | - _broadcast(batch['labels']) | ||
| 100 | - _broadcast(batch['loss_mask']) | ||
| 101 | - _broadcast(batch['attention_mask']) | ||
| 102 | - _broadcast(batch['position_ids']) | ||
| 103 | - | ||
| 104 | - elif mpu.is_pipeline_first_stage(): | ||
| 105 | - _broadcast(batch['tokens']) | ||
| 106 | - _broadcast(batch['attention_mask']) | ||
| 107 | - _broadcast(batch['position_ids']) | ||
| 108 | - if args.schedules_method == 'dualpipev': | ||
| 109 | - _broadcast(batch['loss_mask']) | ||
| 110 | - _broadcast(batch['labels']) | ||
| 111 | - | ||
| 112 | - elif mpu.is_pipeline_last_stage(): | ||
| 113 | - # Multi-Token Prediction (MTP) layers need tokens and position_ids to calculate embedding. | ||
| 114 | - # Currently the Multi-Token Prediction (MTP) layers is fixed on the last stage, so we need | ||
| 115 | - # to broadcast tokens and position_ids to all of the tensor parallel ranks on the last stage. | ||
| 116 | - if args.mtp_num_layers or args.schedules_method == 'dualpipev': | ||
| 117 | - _broadcast(batch['tokens']) | ||
| 118 | - _broadcast(batch['labels']) | ||
| 119 | - _broadcast(batch['loss_mask']) | ||
| 120 | - _broadcast(batch['attention_mask']) | ||
| 121 | - if args.reset_position_ids or args.mtp_num_layers or args.schedules_method == 'dualpipev': | ||
| 122 | - _broadcast(batch['position_ids']) | ||
| 123 | - else: | ||
| 124 | - _broadcast(batch['attention_mask']) | ||
| 125 | - if args.reset_position_ids: | ||
| 126 | - _broadcast(batch['position_ids']) | ||
| 127 | - | ||
| 128 | - else: | ||
| 129 | - if args.enable_share_memory and shm_manager is not None: | ||
| 130 | - actual_seq_len = shm_manager.read() | ||
| 131 | - if '910B' not in acl.get_soc_name() and args.mtp_num_layers and get_post_process_flag(): | ||
| 132 | - set_mtp_position_ids((None, shm_manager)) | ||
| 133 | - | ||
| 134 | - tokens = torch.empty((args.micro_batch_size, args.seq_length), | ||
| 135 | - dtype=torch.int64, | ||
| 136 | - device=torch.cuda.current_device()) | ||
| 137 | - labels = torch.empty((args.micro_batch_size, args.seq_length), | ||
| 138 | - dtype=torch.int64, | ||
| 139 | - device=torch.cuda.current_device()) | ||
| 140 | - loss_mask = torch.empty((args.micro_batch_size, args.seq_length), | ||
| 141 | - dtype=torch.float32, | ||
| 142 | - device=torch.cuda.current_device()) | ||
| 143 | - if args.create_attention_mask_in_dataloader: | ||
| 144 | - attention_mask = torch.empty( | ||
| 145 | - (args.micro_batch_size, 1, args.seq_length, | ||
| 146 | - args.seq_length), dtype=torch.bool, | ||
| 147 | - device=torch.cuda.current_device() | ||
| 148 | - ) | ||
| 149 | - else: | ||
| 150 | - attention_mask = None | ||
| 151 | - position_ids = torch.empty((args.micro_batch_size, args.seq_length), | ||
| 152 | - dtype=torch.int64, | ||
| 153 | - device=torch.cuda.current_device()) | ||
| 154 | - | ||
| 155 | - if args.pipeline_model_parallel_size == 1: | ||
| 156 | - _broadcast(tokens) | ||
| 157 | - _broadcast(labels) | ||
| 158 | - _broadcast(loss_mask) | ||
| 159 | - _broadcast(attention_mask) | ||
| 160 | - _broadcast(position_ids) | ||
| 161 | - | ||
| 162 | - elif mpu.is_pipeline_first_stage(): | ||
| 163 | - _broadcast(tokens) | ||
| 164 | - _broadcast(attention_mask) | ||
| 165 | - _broadcast(position_ids) | ||
| 166 | - if args.schedules_method == 'dualpipev': | ||
| 167 | - _broadcast(loss_mask) | ||
| 168 | - _broadcast(labels) | ||
| 169 | - else: | ||
| 170 | - labels = None | ||
| 171 | - loss_mask = None | ||
| 172 | - | ||
| 173 | - elif mpu.is_pipeline_last_stage(): | ||
| 174 | - if args.mtp_num_layers or args.schedules_method == 'dualpipev': | ||
| 175 | - _broadcast(tokens) | ||
| 176 | - else: | ||
| 177 | - tokens = None | ||
| 178 | - _broadcast(labels) | ||
| 179 | - _broadcast(loss_mask) | ||
| 180 | - _broadcast(attention_mask) | ||
| 181 | - if args.reset_position_ids or args.mtp_num_layers or args.schedules_method == 'dualpipev': | ||
| 182 | - _broadcast(position_ids) | ||
| 183 | - else: | ||
| 184 | - position_ids = None | ||
| 185 | - | ||
| 186 | - else: | ||
| 187 | - tokens = None | ||
| 188 | - labels = None | ||
| 189 | - loss_mask = None | ||
| 190 | - _broadcast(attention_mask) | ||
| 191 | - if args.reset_position_ids: | ||
| 192 | - _broadcast(position_ids) | ||
| 193 | - else: | ||
| 194 | - position_ids = None | ||
| 195 | - | ||
| 196 | - batch = { | ||
| 197 | - 'tokens': tokens, | ||
| 198 | - 'labels': labels, | ||
| 199 | - 'loss_mask': loss_mask, | ||
| 200 | - 'attention_mask': attention_mask, | ||
| 201 | - 'position_ids': position_ids | ||
| 202 | - } | ||
| 203 | - | ||
| 204 | - return batch | ||
原本patch主要用于隐式拷贝优化.cuda(non_blocking=True)性能,需确认下patch原因是否合理