diff --git a/verl/models/mcore/__init__.py b/verl/models/mcore/__init__.py index a0f6e76f3f8..acfdcae46a5 100644 --- a/verl/models/mcore/__init__.py +++ b/verl/models/mcore/__init__.py @@ -16,6 +16,7 @@ from .registry import ( get_mcore_forward_fn, get_mcore_forward_fused_fn, + get_mcore_forward_fused_no_padding_fn, get_mcore_forward_no_padding_fn, get_mcore_weight_converter, hf_to_mcore_config, @@ -28,5 +29,6 @@ "get_mcore_forward_fn", "get_mcore_weight_converter", "get_mcore_forward_fused_fn", + "get_mcore_forward_fused_no_padding_fn", "get_mcore_forward_no_padding_fn", ] diff --git a/verl/models/mcore/model_forward_fused.py b/verl/models/mcore/model_forward_fused.py index 0826caa9c72..70b2b2e659e 100644 --- a/verl/models/mcore/model_forward_fused.py +++ b/verl/models/mcore/model_forward_fused.py @@ -29,12 +29,12 @@ from packaging import version from torch import Tensor -from verl.models.mcore.util import preprocess_packed_seqs +from verl.models.mcore.util import preprocess_packed_seqs, preprocess_thd_no_padding from verl.utils.kernel.linear_cross_entropy import linear_cross_entropy from verl.utils.megatron_utils import unwrap_model from verl.utils.model import CausalLMOutputForPPO -from .util import postprocess_packed_seqs_for_dict_output +from .util import postprocess_packed_seqs_for_dict_output, postprocess_thd_no_padding def _get_patching_model(model: torch.nn.Module): @@ -137,6 +137,79 @@ def fused_forward_model( return fused_forward_model +def fused_forward_no_padding_gen(vision_model: bool = False): + def fused_forward_no_padding( + model, + input_ids: Tensor, + labels: Tensor, + multi_modal_inputs: dict, + temperature: float, + calculate_entropy: bool, + pad_token_id: int, + ): + pre_process = unwrap_model(model).pre_process + post_process = unwrap_model(model).post_process + + input_ids_rmpad, packed_seq_params = preprocess_thd_no_padding(input_ids, pre_process=pre_process) + input_ids_rmpad = input_ids_rmpad.contiguous() + + model_kwargs = {} + if "pixel_values" in multi_modal_inputs: + model_kwargs["pixel_values"] = multi_modal_inputs["pixel_values"].to(input_ids.device) + if "image_grid_thw" in multi_modal_inputs: + model_kwargs["image_grid_thw"] = multi_modal_inputs["image_grid_thw"].to(input_ids.device) + if "pixel_values_videos" in multi_modal_inputs: + model_kwargs["pixel_values_videos"] = multi_modal_inputs["pixel_values_videos"].to(input_ids.device) + if "video_grid_thw" in multi_modal_inputs: + model_kwargs["video_grid_thw"] = multi_modal_inputs["video_grid_thw"].to(input_ids.device) + + attention_mask = None + if vision_model: + input_ids_rmpad = input_ids.to_padded_tensor(pad_token_id) + seqlens_in_batch = input_ids.offsets().diff().to(input_ids.device) + max_seq_len = input_ids_rmpad.shape[1] + attention_mask = torch.arange(max_seq_len, device=input_ids.device).unsqueeze( + 0 + ) < seqlens_in_batch.unsqueeze(1) + + labels_rmpad, _ = preprocess_thd_no_padding(labels, pre_process=True, need_roll=True) + labels_rmpad = labels_rmpad.contiguous() + output_orig: CausalLMOutputForPPO = model( + input_ids=input_ids_rmpad, + attention_mask=attention_mask, + position_ids=None, + packed_seq_params=packed_seq_params, + labels=labels_rmpad, + temperature=temperature, + **model_kwargs, + ) + + if not post_process: + return output_orig + + log_probs = output_orig.log_probs + if log_probs.dim() == 1: + log_probs = log_probs.unsqueeze(0) + log_probs = postprocess_thd_no_padding( + log_probs, packed_seq_params, input_ids, input_ids.shape[0], post_process=post_process + ) + + output = {"log_probs": log_probs} + + if calculate_entropy: + entropy = output_orig.entropy + if entropy.dim() == 1: + entropy = entropy.unsqueeze(0) + entropy = postprocess_thd_no_padding( + entropy, packed_seq_params, input_ids, input_ids.shape[0], post_process=post_process + ) + output["entropy"] = entropy + + return output + + return fused_forward_no_padding + + def _fused_GPTModel_forward( model, input_ids: Tensor, diff --git a/verl/models/mcore/registry.py b/verl/models/mcore/registry.py index bc4679666a0..b1b5c03406b 100644 --- a/verl/models/mcore/registry.py +++ b/verl/models/mcore/registry.py @@ -23,7 +23,7 @@ import torch.nn as nn from .model_forward import gptmodel_forward_no_padding, model_forward_gen -from .model_forward_fused import fused_forward_model_gen +from .model_forward_fused import fused_forward_model_gen, fused_forward_no_padding_gen class SupportedVLM(Enum): @@ -67,6 +67,18 @@ def get_mcore_forward_fused_fn(hf_config) -> Callable: return fused_forward_model_gen(False) +def get_mcore_forward_fused_no_padding_fn(hf_config) -> Callable: + """ + Get the fused forward function for no-padding inputs. + """ + assert len(hf_config.architectures) == 1, "Only one architecture is supported for now" + if hf_config.architectures[0] in supported_vlm: + return fused_forward_no_padding_gen(True) + else: + # default to language model + return fused_forward_no_padding_gen(False) + + # ruff: noqa ######################################################## diff --git a/verl/workers/engine/megatron/transformer_impl.py b/verl/workers/engine/megatron/transformer_impl.py index 90f88c5c154..ec569bcf9b5 100644 --- a/verl/workers/engine/megatron/transformer_impl.py +++ b/verl/workers/engine/megatron/transformer_impl.py @@ -25,7 +25,7 @@ from tensordict import TensorDict import verl.utils.torch_functional as verl_F -from verl.models.mcore import get_mcore_weight_converter +from verl.models.mcore import get_mcore_forward_fused_no_padding_fn, get_mcore_weight_converter from verl.trainer.config import CheckpointConfig from verl.utils import tensordict_utils as tu from verl.utils.checkpoint.megatron_checkpoint_manager import MegatronCheckpointManager @@ -234,6 +234,22 @@ def _build_megatron_module(self): return module + def _maybe_enable_fused_kernels(self): + if not self.engine_config.use_fused_kernels: + return + + if self.is_value_model or self.model_config.mtp.enable: + logger.warning_once( + "Fused kernels are not supported for value models or when MTP is enabled in Megatron engine; disabling." + ) + self.engine_config.use_fused_kernels = False + return + + from verl.models.mcore.model_forward_fused import patch_fused_forward + + for model in self.module: + patch_fused_forward(model) + def _build_optimizer(self): from verl.utils.megatron.optimizer import get_megatron_optimizer, init_megatron_optim_config @@ -274,6 +290,8 @@ def initialize(self): self.module = self._build_megatron_module() + self._maybe_enable_fused_kernels() + if self.model_config.mtp.enable: patch_engine_mtp(self.module, self.model_config) @@ -650,13 +668,6 @@ def forward_step(self, batch_iter: Iterator[TensorDict], model, postprocess_micr multi_modal_inputs = model_inputs["multi_modal_inputs"] loss_mask = model_inputs["loss_mask"] - if not isinstance(temperature, torch.Tensor): - temperature = torch.tensor([temperature] * input_ids.shape[0], device=input_ids.device) - - temperature = temperature.to(torch.float32) - assert temperature.shape[0] == input_ids.shape[0] - temperature = verl_F.expand_as_nested(temperature, input_ids) # (bsz, j1) - if pad_mode == DatasetPadMode.NO_PADDING: label = input_ids.clone() else: @@ -665,7 +676,41 @@ def forward_step(self, batch_iter: Iterator[TensorDict], model, postprocess_micr from verl.models.mcore import get_mcore_forward_no_padding_fn if use_fused_kernels: - raise NotImplementedError("Fused kernels are not supported for megatron engine") + if not self.engine_config.use_remove_padding: + logger.warning_once( + "Fused kernels require `use_remove_padding=True` for Megatron engine. Falling back to non-fused." + ) + use_fused_kernels = False + elif isinstance(temperature, torch.Tensor): + if temperature.numel() != 1: + logger.warning_once( + "Fused kernels do not support per-sample temperature. Falling back to non-fused." + ) + use_fused_kernels = False + else: + temperature_value = float(temperature.item()) + else: + temperature_value = float(temperature) + + if use_fused_kernels: + fused_forward_fn = get_mcore_forward_fused_no_padding_fn(self.model_config.hf_config) + output = fused_forward_fn( + model=model, + input_ids=input_ids, + labels=label, + multi_modal_inputs=multi_modal_inputs, + temperature=temperature_value, + calculate_entropy=calculate_entropy, + pad_token_id=self.model_config.tokenizer.pad_token_id, + ) + return output, partial(postprocess_micro_batch_func, data=batch) + + if not isinstance(temperature, torch.Tensor): + temperature = torch.tensor([temperature] * input_ids.shape[0], device=input_ids.device) + + temperature = temperature.to(torch.float32) + assert temperature.shape[0] == input_ids.shape[0] + temperature = verl_F.expand_as_nested(temperature, input_ids) # (bsz, j1) forward_fn = get_mcore_forward_no_padding_fn(self.model_config.hf_config) diff --git a/verl/workers/engine_workers.py b/verl/workers/engine_workers.py index a3df02c771a..15ba8402cbc 100644 --- a/verl/workers/engine_workers.py +++ b/verl/workers/engine_workers.py @@ -41,12 +41,7 @@ from verl.utils.py_functional import append_to_dict from verl.utils.tensordict_utils import maybe_fix_3d_position_ids from verl.utils.torch_functional import allgather_dict_into_dict -from verl.workers.config import ( - ActorConfig, - HFModelConfig, - RolloutConfig, - TrainingWorkerConfig, -) +from verl.workers.config import ActorConfig, HFModelConfig, RolloutConfig, TrainingWorkerConfig from verl.workers.rollout.base import BaseRollout, get_rollout_class from verl.workers.utils.losses import ppo_loss @@ -88,7 +83,9 @@ def __init__(self, config: TrainingWorkerConfig): ) # we use the one defined in model + # TODO: this is not elegant and should refactor later self.engine_config.use_remove_padding = self.model_config.use_remove_padding + self.engine_config.use_fused_kernels = self.model_config.use_fused_kernels if repatch is not None: # NPU MindSpeed patch, will be refactored with MindSpeedEngine.