已合并
【多模态】【feat.】adapt multimodal models #358
htwang创建于 2025年6月11日
【多模态】【feat.】adapt multimodal models #358
已合并
从refs/pull/358/head合入到master
共 13 个文件变更+165-49
| @@ -15,7 +15,11 @@ class BaseConfig: | |||
| 15 | Method to update parameters from a config dictionary | 15 | 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 default | 205 | 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 = False | 370 | self.swap_attention = False |
| 370 | self.ai_framework = None | 371 | self.ai_framework = None |
| 371 | self.noop_layers = None | 372 | self.noop_layers = None |
| 373 | + | ||
| 372 | self.use_ascend_coc = False | 374 | self.use_ascend_coc = False |
| 373 | self.coc_mode = -1 | 375 | self.coc_mode = -1 |
| 374 | self.coc_parallel_num = 1 | 376 | self.coc_parallel_num = 1 |
| 375 | self.coc_fused_kernel = False | 377 | 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_dict | 76 | 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=sampler | 90 | + 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_manager | 68 | self.sharding_manager = sharding_manager |
| 69 | 69 | ||
| 70 | 70 | ||
| 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 responses | 73 | return responses |
| 74 | 74 | ||
| 75 | 75 | ||
| @@ -6,17 +6,17 @@ from functools import partial | |||
| 6 | import itertools | 6 | import itertools |
| 7 | 7 | ||
| 8 | import torch | 8 | import torch |
| 9 | -from torch.utils.data import DataLoader | 9 | +import torch.nn.functional as F |
| 10 | 10 | ||
| 11 | from mindspeed_rl.models.loss.base_loss_func import BaseLossFunc | 11 | from mindspeed_rl.models.loss.base_loss_func import BaseLossFunc |
| 12 | from mindspeed_rl.models.loss.loss_func_factory import LossFuncFactory | 12 | from mindspeed_rl.models.loss.loss_func_factory import LossFuncFactory |
| 13 | from mindspeed_rl.utils.utils import ( | 13 | from mindspeed_rl.utils.utils import ( |
| 14 | - append_to_dict, generate_mask, generate_position_ids, get_tune_attention_mask | 14 | + append_to_dict, generate_mask, generate_position_ids, get_tune_attention_mask, is_multimodal |
| 15 | ) | 15 | ) |
| 16 | from mindspeed_rl.utils.seqlen_balancing import rearrange_micro_batches, get_reverse_idx | 16 | from mindspeed_rl.utils.seqlen_balancing import rearrange_micro_batches, get_reverse_idx |
| 17 | from mindspeed_rl.utils.remove_padding import preprocess_packed_seqs, postprocess_packed_seqs | 17 | from mindspeed_rl.utils.remove_padding import preprocess_packed_seqs, postprocess_packed_seqs |
| 18 | from mindspeed_rl.utils.compute import get_parallel_state | 18 | from mindspeed_rl.utils.compute import get_parallel_state |
| 19 | -from mindspeed_rl.utils.utils import get_batch_on_this_cp_rank | 19 | +from mindspeed_rl.utils.utils import get_batch_on_this_cp_rank, is_multimodal |
| 20 | 20 | ||
| 21 | 21 | ||
| 22 | class BaseTrainingEngine(ABC): | 22 | class BaseTrainingEngine(ABC): |
| @@ -91,18 +91,26 @@ class BaseTrainingEngine(ABC): | |||
| 91 | self.kwargs = kwargs | 91 | self.kwargs = kwargs |
| 92 | 92 | ||
| 93 | 93 | ||
| 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] = tensor | 108 | + 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 batches | 112 | return batches |
| 105 | - | 113 | + |
| 106 | 114 | ||
| 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 outside | 197 | # 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, batch | 238 | 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_size | 332 | 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 | ||
| 3 | from .vllm_engine import VLLMInferEngine | 3 | from .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 | InferParallelConfig | 38 | InferParallelConfig |
| 39 | ) | 39 | ) |
| 40 | -from mindspeed_rl.utils import get_tokenizer | 40 | +from mindspeed_rl.utils import get_tokenizer, is_multimodal |
| 41 | + | ||
| 41 | 42 | ||
| 42 | logger = Loggers("vllm_engine") | 43 | logger = 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 = None | 274 | self.llm.llm_engine.model_executor.driver_worker.worker.cache_engine = None |
| 273 | self.llm.llm_engine.model_executor.driver_worker.worker.gpu_cache = None | 275 | 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.impl | 278 | 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 = None | 280 | attn_impl.key_cache = None |
| 279 | attn_impl.value_cache = None | 281 | 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.impl | 299 | 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 | 328 | ||
| 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=False | 345 | use_tqdm=False |
| @@ -7,11 +7,12 @@ from .metrics import Metric | |||
| 7 | from .math_eval_toolkit import extract_answer, choice_answer_clean, math_equal | 7 | from .math_eval_toolkit import extract_answer, choice_answer_clean, math_equal |
| 8 | from .utils import ( | 8 | from .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_mask | 10 | + 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_tensor | 73 | 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 | ||
| 5 | import numpy as np | 5 | import numpy as np |
| 6 | import torch | 6 | import 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: int | 85 | + 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 padding | 109 | # 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_new | 124 | 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 clipped | 35 | 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 mask | 48 | 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) |
| 51 | 52 | ||
| 52 | 53 | ||
| 53 | def masked_var(values, mask, unbiased=True): | 54 | def 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 | ||
| 19 | class TestMaskedMean(DistributedTest): | 19 | class TestMaskedMean(DistributedTest): |
| 20 | world_size = 1 | 20 | 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_mean | 23 | 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.0 | 27 | assert result == 2.0 |
| 28 | 28 | ||
| 29 | 29 | ||
| 30 | class TestMaskedVar(DistributedTest): | 30 | class TestMaskedVar(DistributedTest): |
| 31 | world_size = 1 | 31 | 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_var | 34 | from mindspeed_rl.utils.torch_functional import masked_var |
| 35 | import numpy as np | 35 | 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.5 | 39 | + assert np.isclose(result, 2.5, atol=1e-5) |
| 40 | 40 | ||
| 41 | 41 | ||
| 42 | class TestMaskedWhiten(DistributedTest): | 42 | class TestMaskedWhiten(DistributedTest): |
| 43 | world_size = 1 | 43 | 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_whiten | 46 | 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]) |


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