Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions verl/models/mcore/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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",
]
77 changes: 75 additions & 2 deletions verl/models/mcore/model_forward_fused.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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,
Expand Down
14 changes: 13 additions & 1 deletion verl/models/mcore/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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

########################################################
Expand Down
63 changes: 54 additions & 9 deletions verl/workers/engine/megatron/transformer_impl.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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:
Expand All @@ -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)

Expand Down
9 changes: 3 additions & 6 deletions verl/workers/engine_workers.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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.
Expand Down
Loading