已合并
【多模态】【feat.】adapt multimodal models #358
htwang创建于 2025年6月11日
【多模态】【feat.】adapt multimodal models #358
已合并
htwang创建于 2025年6月11日
从refs/pull/358/head合入到master
共 13 个文件变更+165-49
@@ -15,7 +15,11 @@ class BaseConfig:
15 Method to update parameters from a config dictionary15 Method to update parameters from a config dictionary
16 '''16 '''
17 if 'model' in config_dict:17 if 'model' in config_dict:
18- self.update(model_config_dict[config_dict['model']])18+ # if str, parsed as file path, used for multi-modal
19+ if isinstance(model_config_dict, str):
20+ self._process_multi_modal(model_config_dict)
21+ else:
22+ self.update(model_config_dict[config_dict['model']])
19 23 
20 for key, value in config_dict.items():24 for key, value in config_dict.items():
21 if key == 'model':25 if key == 'model':
@@ -37,3 +41,20 @@ class BaseConfig:
37 41 
38 def dict(self):42 def dict(self):
39 return self.__dict__43 return self.__dict__
44+ 
45+ def _process_multi_modal(self, model_config_dict):
46+ import json
47+ from pathlib import Path
48+ config_path = Path(model_config_dict)
49+ if not config_path.is_file():
50+ raise FileNotFoundError(f'model json file: {str(model_config_dict)} is not found!')
51+ else:
52+ config = json.loads(config_path.read_text())
53+ # used for actor_hybrid_worker initialize megatron resharding manager
54+ img_pp_layers = config.get('image_encoder', {}).get('vision_encoder', {}).get('pipeline_num_layers', None)
55+ llm_pp_layers = config.get('text_decoder', {}).get('pipeline_num_layers', None)
56+ if img_pp_layers is None or llm_pp_layers is None:
57+ raise ValueError(f'`pipeline_num_layers` should be set in config file: {str(model_config_dict)}.')
58+ setattr(self, 'num_layer_list', [img_pp_layers, llm_pp_layers])
59+ # used for mindspeed mm
60+ setattr(self, "mm_model", model_config_dict)
@@ -109,7 +109,7 @@ class MegatronConfig(BaseConfig):
109 bf16: Whether to use BF16 (default: False)109 bf16: Whether to use BF16 (default: False)
110 use_distributed_optimizer: Use distributed optimizer (default: False)110 use_distributed_optimizer: Use distributed optimizer (default: False)
111 is_instruction_dataset: Whether the dataset is instruction-based (default: False)111 is_instruction_dataset: Whether the dataset is instruction-based (default: False)
112- is_pairwise_dataset: Whether the dataset is pairwise format that has a chosen sequence and rejected 112+ is_pairwise_dataset: Whether the dataset is pairwise format that has a chosen sequence and rejected
113 sequence, which usually used in reinforce learning (default: False)113 sequence, which usually used in reinforce learning (default: False)
114 variable_seq_lengths: Whether to use variable sequence lengths (default: False)114 variable_seq_lengths: Whether to use variable sequence lengths (default: False)
115 no_shuffle: Whether to shuffle the dataset (default: False)115 no_shuffle: Whether to shuffle the dataset (default: False)
@@ -153,7 +153,7 @@ class MegatronConfig(BaseConfig):
153 eval_interval: Interval between running evaluation on validation set (default: 1000)153 eval_interval: Interval between running evaluation on validation set (default: 1000)
154 seed: Random seed used for python, numpy, pytorch, and cuda (default: 1234)154 seed: Random seed used for python, numpy, pytorch, and cuda (default: 1234)
155 vocab_extra_ids: Number of additional vocabulary tokens. They are used for span masking in the T5 model (default: 0)155 vocab_extra_ids: Number of additional vocabulary tokens. They are used for span masking in the T5 model (default: 0)
156- use_tp_pp_dp_mapping: If set, distributed ranks initialize order is changed from tp-dp-pp to tp-pp-dp. 156+ use_tp_pp_dp_mapping: If set, distributed ranks initialize order is changed from tp-dp-pp to tp-pp-dp.
157 Make sure EP and CP aren't used with this option enabled with this option enabled (default: False)157 Make sure EP and CP aren't used with this option enabled with this option enabled (default: False)
158 log_interval: Report loss and timing interval (default: 100)158 log_interval: Report loss and timing interval (default: 100)
159 load_checkpoint_loosely: Enable loading checkpoint not strictly (default: False)159 load_checkpoint_loosely: Enable loading checkpoint not strictly (default: False)
@@ -205,6 +205,7 @@ class MegatronConfig(BaseConfig):
205 coc_mode: 0=original, 1=rewrite, 2=coc default205 coc_mode: 0=original, 1=rewrite, 2=coc default
206 coc_parallel_num: number of parallel in CoC features (default: 1)206 coc_parallel_num: number of parallel in CoC features (default: 1)
207 coc_fused_kernel: switch to use fused kernel in CoC (default: False)207 coc_fused_kernel: switch to use fused kernel in CoC (default: False)
208+ mm_model: config for multimodal models
208 '''209 '''
209 210 
210 def __init__(self, training_config: Dict, model_config: Dict):211 def __init__(self, training_config: Dict, model_config: Dict):
@@ -369,9 +370,13 @@ class MegatronConfig(BaseConfig):
369 self.swap_attention = False370 self.swap_attention = False
370 self.ai_framework = None371 self.ai_framework = None
371 self.noop_layers = None372 self.noop_layers = None
373+ 
372 self.use_ascend_coc = False374 self.use_ascend_coc = False
373 self.coc_mode = -1375 self.coc_mode = -1
374 self.coc_parallel_num = 1376 self.coc_parallel_num = 1
375 self.coc_fused_kernel = False377 self.coc_fused_kernel = False
376- 378+ 
379+ # used for multimodal models
380+ self.mm_model = None
381+ 
377 self.update(training_config, model_config)382 self.update(training_config, model_config)
@@ -74,7 +74,7 @@ class MultiModalDataLoader(torch.utils.data.DataLoader):
74 batch_dict['prompts'] = [torch.tensor(i) for i in batch_dict['prompts']]74 batch_dict['prompts'] = [torch.tensor(i) for i in batch_dict['prompts']]
75 75 
76 return batch_dict76 return batch_dict
77- 77+ 
78 if not no_shuffle:78 if not no_shuffle:
79 train_dataloader_generator = torch.Generator()79 train_dataloader_generator = torch.Generator()
80 train_dataloader_generator.manual_seed(seed)80 train_dataloader_generator.manual_seed(seed)
@@ -83,9 +83,10 @@ class MultiModalDataLoader(torch.utils.data.DataLoader):
83 sampler = SequentialSampler(data_source=dataset)83 sampler = SequentialSampler(data_source=dataset)
84 84 
85 super().__init__(dataset,85 super().__init__(dataset,
86- batch_size=global_batch_size,86+ num_workers=num_workers,
87- num_workers=num_workers,87+ generator=torch.Generator().manual_seed(seed),
88- drop_last=True,88+ collate_fn=collate_fn,
89- collate_fn=collate_fn,89+ pin_memory=True,
90- sampler=sampler90+ sampler=sampler,
91- )91+ batch_size=global_batch_size,
92+ drop_last=True)
@@ -68,8 +68,8 @@ class ActorRolloutHybrid(ABC):
68 self.sharding_manager = sharding_manager68 self.sharding_manager = sharding_manager
69 69 
70 @mstx_timer_decorator70 @mstx_timer_decorator
71- def generate_sequences(self, prompts_list: List[List[int]]) -> Tensor:71+ def generate_sequences(self, prompts_list: List[List[int]], **kwargs) -> Tensor:
72- responses = self.inference_actor.generate_sequences(prompts_list)[0]72+ responses = self.inference_actor.generate_sequences(prompts_list, **kwargs)[0]
73 return responses73 return responses
74 74 
75 @mstx_timer_decorator75 @mstx_timer_decorator
@@ -6,17 +6,17 @@ from functools import partial
6import itertools6import itertools
7 7 
8import torch8import torch
9-from torch.utils.data import DataLoader9+import torch.nn.functional as F
10 10 
11from mindspeed_rl.models.loss.base_loss_func import BaseLossFunc11from mindspeed_rl.models.loss.base_loss_func import BaseLossFunc
12from mindspeed_rl.models.loss.loss_func_factory import LossFuncFactory12from mindspeed_rl.models.loss.loss_func_factory import LossFuncFactory
13from mindspeed_rl.utils.utils import (13from mindspeed_rl.utils.utils import (
14- append_to_dict, generate_mask, generate_position_ids, get_tune_attention_mask14+ append_to_dict, generate_mask, generate_position_ids, get_tune_attention_mask, is_multimodal
15)15)
16from mindspeed_rl.utils.seqlen_balancing import rearrange_micro_batches, get_reverse_idx16from mindspeed_rl.utils.seqlen_balancing import rearrange_micro_batches, get_reverse_idx
17from mindspeed_rl.utils.remove_padding import preprocess_packed_seqs, postprocess_packed_seqs17from mindspeed_rl.utils.remove_padding import preprocess_packed_seqs, postprocess_packed_seqs
18from mindspeed_rl.utils.compute import get_parallel_state18from mindspeed_rl.utils.compute import get_parallel_state
19-from mindspeed_rl.utils.utils import get_batch_on_this_cp_rank19+from mindspeed_rl.utils.utils import get_batch_on_this_cp_rank, is_multimodal
20 20 
21 21 
22class BaseTrainingEngine(ABC):22class BaseTrainingEngine(ABC):
@@ -91,18 +91,26 @@ class BaseTrainingEngine(ABC):
91 self.kwargs = kwargs91 self.kwargs = kwargs
92 92 
93 @staticmethod93 @staticmethod
94- def _split_batches(batch: Dict, batch_size: int, shuffle_mini_batch: bool, dim: int = 0) -> List[Dict]:94+ def _split_batches(batch: Dict, batch_size: int, shuffle_mini_batch: bool, dim: int = 0, keep_list: bool = False) -> List[Dict]:
95 batches = []95 batches = []
96- for key, tensors in batch.items():96+ for key, value in batch.items():
97- for index, tensor in enumerate(torch.split(tensors, batch_size, dim)):97+ if isinstance(value, torch.Tensor):
98+ split_values = torch.split(value, batch_size, dim)
99+ elif isinstance(value, List): # 多模态场景
100+ num_batches = (len(value) + batch_size - 1) // batch_size
101+ if keep_list: # 保持值为list类型
102+ split_values = [value[i * batch_size: (i + 1) * batch_size] for i in range(num_batches)]
103+ else:
104+ split_values = [torch.concat(value[i * batch_size: (i + 1) * batch_size]) for i in range(num_batches)]
105+ for index, split_value in enumerate(split_values):
98 if index >= len(batches):106 if index >= len(batches):
99 batches.append({})107 batches.append({})
100- batches[index][key] = tensor108+ batches[index][key] = split_value
101 109 
102 if shuffle_mini_batch:110 if shuffle_mini_batch:
103 random.shuffle(batches)111 random.shuffle(batches)
104 return batches112 return batches
105- 113+ 
106 @staticmethod114 @staticmethod
107 def _split_batches_with_dynamic_bsz(batch: Dict, max_packing_token: int) -> (List[Dict], List[List[int]]):115 def _split_batches_with_dynamic_bsz(batch: Dict, max_packing_token: int) -> (List[Dict], List[List[int]]):
108 seq_len_list = []116 seq_len_list = []
@@ -134,7 +142,17 @@ class BaseTrainingEngine(ABC):
134 142 
135 def forward_step(batch_iter, model):143 def forward_step(batch_iter, model):
136 cp_size = get_parallel_state().get_context_parallel_world_size()144 cp_size = get_parallel_state().get_context_parallel_world_size()
137- if self.use_remove_padding:145+ if is_multimodal():
146+ process_batch, seqlens_in_batch, cu_seqlens_padded = self._get_mulitmodal_forward_batch_info(batch_iter)
147+ output = model(**process_batch)
148+ if post_process:
149+ output = postprocess_packed_seqs(output=output['logits'],
150+ seqlens_in_batch=seqlens_in_batch,
151+ cu_seqlens_padded=cu_seqlens_padded,
152+ seq_len=seq_len,
153+ return_tensor=True)
154+ output.div_(self.temperature)
155+ elif self.use_remove_padding:
138 input_ids, position_ids, process_batch, seqlens_in_batch, cu_seqlens_padded = self._get_forward_batch_info(batch_iter)156 input_ids, position_ids, process_batch, seqlens_in_batch, cu_seqlens_padded = self._get_forward_batch_info(batch_iter)
139 self.set_actual_seq_len(cu_seqlens_padded.tolist())157 self.set_actual_seq_len(cu_seqlens_padded.tolist())
140 output_orig = model(input_ids=input_ids, attention_mask=None, position_ids=position_ids)158 output_orig = model(input_ids=input_ids, attention_mask=None, position_ids=position_ids)
@@ -175,7 +193,7 @@ class BaseTrainingEngine(ABC):
175 forward_only=forward_only,193 forward_only=forward_only,
176 collect_non_loss_data=forward_only,194 collect_non_loss_data=forward_only,
177 )195 )
178- 196+ 
179 # Reverse the batch index to be the same outside197 # Reverse the batch index to be the same outside
180 if self.use_dynamic_bsz and forward_only and post_process:198 if self.use_dynamic_bsz and forward_only and post_process:
181 losses_reduced_list = torch.cat(losses_reduced, dim=0)199 losses_reduced_list = torch.cat(losses_reduced, dim=0)
@@ -219,6 +237,41 @@ class BaseTrainingEngine(ABC):
219 237 
220 return input_ids, attention_mask, position_ids, batch238 return input_ids, attention_mask, position_ids, batch
221 239 
240+ def _get_mulitmodal_forward_batch_info(self, batch_iter):
241+ batch = next(batch_iter)
242+ input_ids = batch['input_ids']
243+ batch_size = input_ids.size(0)
244+ 
245+ response_attention_mask = generate_mask(batch['responses'], batch['response_length']).to(input_ids.device)
246+ attention_mask = torch.cat((batch['attention_mask'], response_attention_mask), dim=-1).bool()
247+ 
248+ delta_position_id = torch.tensor(generate_position_ids(batch['responses'])).to(input_ids.device) + 1
249+ delta_position_id = delta_position_id.unsqueeze(1).expand(batch_size, 1, -1)
250+ 
251+ if batch['position_ids'].dim() == 3: # qwen2vl mrope
252+ delta_position_id = delta_position_id.expand(batch_size, 3, -1)
253+ else:
254+ batch['position_ids'] = batch['position_ids'].view(batch_size, 1, -1)
255+ response_position_ids = batch['position_ids'][..., -1:] + delta_position_id
256+ position_ids = torch.cat([batch['position_ids'], response_position_ids], dim=-1)
257+ 
258+ input_ids_rmpad_list = []
259+ position_ids_rmpad_list = []
260+ for i in range(batch_size):
261+ input_ids_rmpad_list.append(input_ids[i].masked_select(attention_mask[i]))
262+ masked_position_ids = position_ids[i].masked_select(attention_mask[i].expand_as(position_ids[i]))
263+ position_ids_rmpad_list.append(masked_position_ids.view(delta_position_id.size(1), -1).t())
264+ input_ids_rmpad = torch.cat(input_ids_rmpad_list).unsqueeze(0)
265+ position_ids_rmpad = torch.cat(position_ids_rmpad_list).t().unsqueeze(1)
266+ 
267+ seqlens_in_batch = attention_mask.sum(dim=-1, dtype=torch.int32)
268+ cu_seqlens_padded = F.pad(torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0))
269+ 
270+ batch['input_ids'] = input_ids_rmpad
271+ batch['position_ids'] = position_ids_rmpad
272+ batch['attention_mask'] = None
273+ return batch, seqlens_in_batch, cu_seqlens_padded
274+ 
222 def post_process_forward_backward_output(self, output: [torch.Tensor],275 def post_process_forward_backward_output(self, output: [torch.Tensor],
223 batch: Dict[str, torch.Tensor]) -> torch.Tensor:276 batch: Dict[str, torch.Tensor]) -> torch.Tensor:
224 """277 """
@@ -236,7 +289,9 @@ class BaseTrainingEngine(ABC):
236 :return: 模型前向计算结果。289 :return: 模型前向计算结果。
237 """290 """
238 for k, v in data.items():291 for k, v in data.items():
239- if v is not None:292+ if isinstance(v, List):
293+ data[k] = [t.to(next(self.model[0].parameters()).device) for t in v]
294+ else:
240 data[k] = v.to(next(self.model[0].parameters()).device)295 data[k] = v.to(next(self.model[0].parameters()).device)
241 for model_module in self.model:296 for model_module in self.model:
242 model_module.eval()297 model_module.eval()
@@ -256,9 +311,12 @@ class BaseTrainingEngine(ABC):
256 grad_norm_list = []311 grad_norm_list = []
257 for k, v in data.items():312 for k, v in data.items():
258 if v is not None:313 if v is not None:
259- data[k] = v.to(next(self.model[0].parameters()).device)314+ if isinstance(v, List):
315+ data[k] = [t.to(next(self.model[0].parameters()).device) for t in v]
316+ else:
317+ data[k] = v.to(next(self.model[0].parameters()).device)
260 mini_batches = self._split_batches(data, batch_size=self.mini_batch_size_per_dp,318 mini_batches = self._split_batches(data, batch_size=self.mini_batch_size_per_dp,
261- shuffle_mini_batch=self.shuffle_mini_batch, dim=0)319+ shuffle_mini_batch=self.shuffle_mini_batch, dim=0, keep_list=True)
262 for model_module in self.model:320 for model_module in self.model:
263 model_module.train()321 model_module.train()
264 for _ in range(self.epochs):322 for _ in range(self.epochs):
@@ -273,7 +331,7 @@ class BaseTrainingEngine(ABC):
273 data_parallel_world_size = get_parallel_state().get_data_parallel_world_size()331 data_parallel_world_size = get_parallel_state().get_data_parallel_world_size()
274 increment = self.mini_batch_size_per_dp * data_parallel_world_size332 increment = self.mini_batch_size_per_dp * data_parallel_world_size
275 self.opt_param_scheduler.step(increment=increment)333 self.opt_param_scheduler.step(increment=increment)
276- grad_norm_list.append(grad_norm) 334+ grad_norm_list.append(grad_norm)
277 335 
278 for metric in metric_micro_batch:336 for metric in metric_micro_batch:
279 append_to_dict(metrics, metric) # append the metric from this micro-batch to global metrics.337 append_to_dict(metrics, metric) # append the metric from this micro-batch to global metrics.
@@ -1,6 +1,7 @@
1# Copyright (c) 2025, HUAWEI CORPORATION. All rights reserved.1# Copyright (c) 2025, HUAWEI CORPORATION. All rights reserved.
2 2 
3from .vllm_engine import VLLMInferEngine3from .vllm_engine import VLLMInferEngine
4+from .vllm_adapter import patch
4 5 
5 6 
6__all__ = ['VLLMInferEngine']7__all__ = ['VLLMInferEngine']
@@ -37,7 +37,8 @@ from mindspeed_rl.models.rollout.vllm_adapter.megatron_weight_loaders import (
37 update_megatron_weight_loader,37 update_megatron_weight_loader,
38 InferParallelConfig38 InferParallelConfig
39)39)
40-from mindspeed_rl.utils import get_tokenizer40+from mindspeed_rl.utils import get_tokenizer, is_multimodal
41+ 
41 42 
42logger = Loggers("vllm_engine")43logger = Loggers("vllm_engine")
43 44 
@@ -179,6 +180,7 @@ class VLLMInferEngine(BaseInferEngine):
179 gpu_memory_utilization=gpu_memory_utilization,180 gpu_memory_utilization=gpu_memory_utilization,
180 max_num_seqs=max_num_seqs,181 max_num_seqs=max_num_seqs,
181 max_model_len=max_model_len,182 max_model_len=max_model_len,
183+ seed=self.sampling_params.seed,
182 additional_config={184 additional_config={
183 'expert_tensor_parallel_size': infer_expert_tensor_parallel_size,185 'expert_tensor_parallel_size': infer_expert_tensor_parallel_size,
184 'enable_graph_mode': int(os.environ.get('VLLM_ENABLE_GRAPH_MODE', '0')),186 'enable_graph_mode': int(os.environ.get('VLLM_ENABLE_GRAPH_MODE', '0')),
@@ -271,12 +273,19 @@ class VLLMInferEngine(BaseInferEngine):
271 else:273 else:
272 self.llm.llm_engine.model_executor.driver_worker.worker.cache_engine = None274 self.llm.llm_engine.model_executor.driver_worker.worker.cache_engine = None
273 self.llm.llm_engine.model_executor.driver_worker.worker.gpu_cache = None275 self.llm.llm_engine.model_executor.driver_worker.worker.gpu_cache = None
274- if hasattr(self.model.model.layers[0].self_attn, "attn"):276+ if hasattr(self.model, 'model') and hasattr(self.model.model.layers[0].self_attn, "attn"):
275 for i in range(self.model.model.start_layer, self.model.model.end_layer):277 for i in range(self.model.model.start_layer, self.model.model.end_layer):
276 attn_impl = self.model.model.layers[i].self_attn.attn.impl278 attn_impl = self.model.model.layers[i].self_attn.attn.impl
277 if hasattr(attn_impl, "key_cache"):279 if hasattr(attn_impl, "key_cache"):
278 attn_impl.key_cache = None280 attn_impl.key_cache = None
279 attn_impl.value_cache = None281 attn_impl.value_cache = None
282+ # 多模态kv cache
283+ elif hasattr(self.model, 'language_model') and hasattr(self.model.language_model.model.layers[0].self_attn, "attn"):
284+ for i in range(self.model.language_model.model.start_layer, self.model.language_model.model.end_layer):
285+ attn_impl = self.model.language_model.model.layers[i].self_attn.attn.impl
286+ if hasattr(attn_impl, "key_cache"):
287+ attn_impl.key_cache = None
288+ attn_impl.value_cache = None
280 289 
281 gc.collect()290 gc.collect()
282 torch.cuda.empty_cache()291 torch.cuda.empty_cache()
@@ -285,7 +294,7 @@ class VLLMInferEngine(BaseInferEngine):
285 def offload_model_weights(self):294 def offload_model_weights(self):
286 for name, params in self.model.named_parameters():295 for name, params in self.model.named_parameters():
287 params.data = self.cpu_model[name]296 params.data = self.cpu_model[name]
288- if hasattr(self.model.model.layers[-1].self_attn, "mla_attn"):297+ if hasattr(self.model, 'model') and hasattr(self.model.model.layers[-1].self_attn, "mla_attn"):
289 for i in range(self.model.model.start_layer, self.model.model.end_layer):298 for i in range(self.model.model.start_layer, self.model.model.end_layer):
290 mla = self.model.model.layers[i].self_attn.mla_attn.impl299 mla = self.model.model.layers[i].self_attn.mla_attn.impl
291 if hasattr(mla, "w_kc"):300 if hasattr(mla, "w_kc"):
@@ -302,7 +311,7 @@ class VLLMInferEngine(BaseInferEngine):
302 self.model,311 self.model,
303 infer_parallel_config,312 infer_parallel_config,
304 self.hf_config)313 self.hf_config)
305- if hasattr(self.model.model.layers[0].self_attn, "mla_attn"):314+ if hasattr(self.model, 'model') and hasattr(self.model.model.layers[0].self_attn, "mla_attn"):
306 self._process_mla()315 self._process_mla()
307 316 
308 def _process_mla(self):317 def _process_mla(self):
@@ -319,9 +328,18 @@ class VLLMInferEngine(BaseInferEngine):
319 @torch.no_grad()328 @torch.no_grad()
320 def generate_sequences(self, idx_list, **kwargs):329 def generate_sequences(self, idx_list, **kwargs):
321 self.init_cache_engine()330 self.init_cache_engine()
331+ if is_multimodal():
332+ images = kwargs.pop("extra_info")
333+ prompts = [
334+ {"prompt_token_ids": prompt, "multi_modal_data": {"image": image}}
335+ for prompt, image in zip(idx_list, images['image'])
336+ ]
337+ idx_list = None
338+ else:
339+ prompts = None
322 with self.update_sampling_params(**kwargs):340 with self.update_sampling_params(**kwargs):
323 response = self.llm.generate(341 response = self.llm.generate(
324- prompts=None,342+ prompts=prompts,
325 sampling_params=self.sampling_params,343 sampling_params=self.sampling_params,
326 prompt_token_ids=idx_list,344 prompt_token_ids=idx_list,
327 use_tqdm=False345 use_tqdm=False
@@ -7,11 +7,12 @@ from .metrics import Metric
7from .math_eval_toolkit import extract_answer, choice_answer_clean, math_equal7from .math_eval_toolkit import extract_answer, choice_answer_clean, math_equal
8from .utils import (8from .utils import (
9 get_batch_metrices_mean, num_floating_point_operations,9 get_batch_metrices_mean, num_floating_point_operations,
10- seed_all, synchronize_time, parse_args_from_config, get_tune_attention_mask10+ seed_all, synchronize_time, parse_args_from_config, get_tune_attention_mask,
11+ is_multimodal
11)12)
12 13 
13__all__ = ['get_tokenizer', 'Loggers', 'WandbLogger', 'Metric',14__all__ = ['get_tokenizer', 'Loggers', 'WandbLogger', 'Metric',
14 'get_batch_metrices_mean', 'num_floating_point_operations',15 'get_batch_metrices_mean', 'num_floating_point_operations',
15 'seed_all', 'synchronize_time', 'parse_args_from_config',16 'seed_all', 'synchronize_time', 'parse_args_from_config',
16 'extract_answer', 'choice_answer_clean', 'math_equal',17 'extract_answer', 'choice_answer_clean', 'math_equal',
17- 'get_tune_attention_mask']18+ 'get_tune_attention_mask', 'is_multimodal']
@@ -73,7 +73,7 @@ def truncate_middle_and_pad(responses, input_tensor, truncate_lengths, pad_value
73 return output_tensor73 return output_tensor
74 74 
75 75 
76-def truncate_rows(tensor, index_tensor):76+def truncate_rows(tensor, index_tensor, left_pad=False):
77 """77 """
78 tensor: 二维 Tensor,形状为 (mbs, seq_len)78 tensor: 二维 Tensor,形状为 (mbs, seq_len)
79 index_tensor: 二维 Tensor,形状为 (mbs, 1),表示每一行截断的位置79 index_tensor: 二维 Tensor,形状为 (mbs, 1),表示每一行截断的位置
@@ -85,7 +85,10 @@ def truncate_rows(tensor, index_tensor):
85 # 获取当前行的截断索引85 # 获取当前行的截断索引
86 trunc_idx = index_tensor[i].item()86 trunc_idx = index_tensor[i].item()
87 # 截断当前行87 # 截断当前行
88- truncated_row = tensor[i, :trunc_idx].cpu()88+ if left_pad:
89+ truncated_row = tensor[i, -trunc_idx:].cpu()
90+ else:
91+ truncated_row = tensor[i, :trunc_idx].cpu()
89 # 将截断后的行添加到列表中92 # 将截断后的行添加到列表中
90 truncated_tensors.append(truncated_row)93 truncated_tensors.append(truncated_row)
91 94 
@@ -1,6 +1,6 @@
1# Copyright (c) 2025, HUAWEI CORPORATION. All rights reserved.1# Copyright (c) 2025, HUAWEI CORPORATION. All rights reserved.
2 2 
3-from typing import Tuple 3+from typing import Tuple
4 4 
5import numpy as np5import numpy as np
6import torch6import torch
@@ -82,7 +82,8 @@ def postprocess_packed_seqs(
82 output: torch.Tensor,82 output: torch.Tensor,
83 seqlens_in_batch: torch.Tensor,83 seqlens_in_batch: torch.Tensor,
84 cu_seqlens_padded: torch.Tensor,84 cu_seqlens_padded: torch.Tensor,
85- seq_len: int85+ seq_len: int,
86+ return_tensor: bool = False
86) -> torch.Tensor:87) -> torch.Tensor:
87 """88 """
88 Unpacks a packed output tensor back into the original batch shape, restoring padding.89 Unpacks a packed output tensor back into the original batch shape, restoring padding.
@@ -107,12 +108,18 @@ def postprocess_packed_seqs(
107 108 
108 # Prepare new output with padding109 # Prepare new output with padding
109 batch_size = seqlens_in_batch.shape[0]110 batch_size = seqlens_in_batch.shape[0]
110- full_shape = [batch_size, seq_len] + list(output.shape[2:])111+ if return_tensor:
111- output_new = []112+ full_shape = [batch_size, seq_len] + list(output.shape[2:])
113+ output_new = torch.zeros(full_shape, dtype=output.dtype, device=output.device)
114+ else:
115+ output_new = []
112 for i in range(batch_size):116 for i in range(batch_size):
113 start = cu_seqlens_padded[i].item()117 start = cu_seqlens_padded[i].item()
114 length = seqlens_in_batch[i].item()118 length = seqlens_in_batch[i].item()
115- output_new.append(output[0, start:start + length])119+ if return_tensor:
120+ output_new[i, :length] = output[0, start:start + length]
121+ else:
122+ output_new.append(output[0, start:start + length])
116 123 
117 return output_new124 return output_new
118 125 
@@ -1,4 +1,5 @@
1# Copyright (c) 2025, HUAWEI CORPORATION. All rights reserved.1# Copyright (c) 2025, HUAWEI CORPORATION. All rights reserved.
2+# Copyright 2024 Bytedance Ltd. and/or its affiliates
2#3#
3# Licensed under the Apache License, Version 2.0 (the "License");4# Licensed under the Apache License, Version 2.0 (the "License");
4# you may not use this file except in compliance with the License.5# you may not use this file except in compliance with the License.
@@ -34,7 +35,7 @@ def clip_by_value(x, tensor_min, tensor_max):
34 return clipped35 return clipped
35 36 
36 37 
37-def masked_mean(values, mask, axis=None):38+def masked_mean(values, mask, axis=None, epsilon=1e-8):
38 """39 """
39 Compute mean of tensor with a masked values.40 Compute mean of tensor with a masked values.
40 41 
@@ -47,7 +48,7 @@ def masked_mean(values, mask, axis=None):
47 The mean of the data after applying the mask48 The mean of the data after applying the mask
48 """49 """
49 50 
50- return (values * mask).sum(axis=axis) / mask.sum(axis=axis)51+ return (values * mask).sum(axis=axis) / (mask.sum(axis=axis) + epsilon)
HH
Hhtwang2025年6月18日

?这里咋加了个,只有多模态用了?

这里其实参考的verl来的,https://github.com/volcengine/verl/blob/main/verl/utils/torch_functional.py#L166 这里主要是为了对精度和verl对的更齐

likedislike
Hhtwang2025年6月18日

?这里咋加了个,只有多模态用了?

这里当前msrl开发时去掉的初衷不太清楚,但其实加这个epsilon更保险点

likedislike
51 52 
52 53 
53def masked_var(values, mask, unbiased=True):54def masked_var(values, mask, unbiased=True):
@@ -270,7 +270,7 @@ class ActorHybridWorkerBase(BaseWorker):
270 prompts = truncate_rows(prompts_data, prompt_length_data)270 prompts = truncate_rows(prompts_data, prompt_length_data)
271 prompts_list = [prompt.numpy().tolist() for prompt in prompts]271 prompts_list = [prompt.numpy().tolist() for prompt in prompts]
272 272 
273- responses_pad_right = self.actor_hybrid.generate_sequences(copy.deepcopy(prompts_list))273+ responses_pad_right = self.actor_hybrid.generate_sequences(copy.deepcopy(prompts_list), extra_info=batch_data)
274 responses = remove_padding_and_split_to_list(responses_pad_right, self.tokenizer.eod, pad_token_id)274 responses = remove_padding_and_split_to_list(responses_pad_right, self.tokenizer.eod, pad_token_id)
275 275 
276 responses_length = [torch.tensor([len(response)]) for response in responses]276 responses_length = [torch.tensor([len(response)]) for response in responses]
@@ -18,30 +18,30 @@ class TestClipByValue(DistributedTest):
18 18 
19class TestMaskedMean(DistributedTest):19class TestMaskedMean(DistributedTest):
20 world_size = 120 world_size = 1
21- 21+ 
22 def test_masked_mean_1d(self):22 def test_masked_mean_1d(self):
23 from mindspeed_rl.utils.torch_functional import masked_mean23 from mindspeed_rl.utils.torch_functional import masked_mean
24 values = torch.tensor([1.0, 2.0, 3.0])24 values = torch.tensor([1.0, 2.0, 3.0])
25 mask = torch.tensor([1.0, 0.0, 1.0])25 mask = torch.tensor([1.0, 0.0, 1.0])
26- result = masked_mean(values, mask)26+ result = masked_mean(values, mask, epsilon=0)
27 assert result == 2.027 assert result == 2.0
28 28 
29 29 
30class TestMaskedVar(DistributedTest):30class TestMaskedVar(DistributedTest):
31 world_size = 131 world_size = 1
32- 32+ 
33 def test_masked_var_unbiased_true(self):33 def test_masked_var_unbiased_true(self):
34 from mindspeed_rl.utils.torch_functional import masked_var34 from mindspeed_rl.utils.torch_functional import masked_var
35 import numpy as np35 import numpy as np
36 values = np.array([1, 2, 3, 4, 5])36 values = np.array([1, 2, 3, 4, 5])
37 mask = np.array([1, 1, 1, 1, 1])37 mask = np.array([1, 1, 1, 1, 1])
38 result = masked_var(values, mask, unbiased=True)38 result = masked_var(values, mask, unbiased=True)
39- assert result == 2.539+ assert np.isclose(result, 2.5, atol=1e-5)
40 40 
41 41 
42class TestMaskedWhiten(DistributedTest):42class TestMaskedWhiten(DistributedTest):
43 world_size = 143 world_size = 1
44- 44+ 
45 def test_masked_whiten_shift_mean_true(self):45 def test_masked_whiten_shift_mean_true(self):
46 from mindspeed_rl.utils.torch_functional import masked_whiten46 from mindspeed_rl.utils.torch_functional import masked_whiten
47 values = torch.tensor([1.0, 2.0, 3.0, 4.0])47 values = torch.tensor([1.0, 2.0, 3.0, 4.0])