已合并
[mindspore][patch][master]del patch #3696
王一博创建于 2025年11月15日
[mindspore][patch][master]del patch #3696
已合并
王一博创建于 2025年11月15日
共 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_scaling85 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 memory149 # share memory
157 if args.enable_share_memory:150 if args.enable_share_memory:
158 from ..mindspore.tasks.dataset.shared_memory_manager import SharedMemoryManager151 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_len154 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 parameter158 # Optimizer: load and save parameter
@@ -257,7 +250,6 @@ def _patch_fused_operators(args):
257 # Matmul add ops250 # Matmul add ops
258 from mindspeed.mindspore.ops.npu_matmul_add import npu_matmul_add_fp32251 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 v2254 # Fused AdamW v2
263 from torch import npu_apply_fused_adamw_v2255 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
8import numpy as np5import 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 
21def _compute_actual_seq_len(origin_seq):8def _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 res17 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"],
ZhihaoLi
ZhihaoLiZhihaoLi2025年11月22日
已过期

原本patch主要用于隐式拷贝优化.cuda(non_blocking=True)性能,需确认下patch原因是否合理

likedislike
王一博
王一博
2025年11月22日 评论:
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