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
137 changes: 137 additions & 0 deletions tests/trainer/verl_compat_test.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,137 @@
"""CPU regression tests for the verl 0.9 worker configuration/input contracts."""
from types import SimpleNamespace
from unittest import mock

import pytest
import torch
from omegaconf import OmegaConf
from tensordict import TensorDict

pytest.importorskip("verl", minversion="0.9.0")

from verl.trainer.config import CheckpointConfig # noqa: E402
from verl.utils.config import omega_conf_to_dataclass # noqa: E402
from verl.workers.config.checkpoint import McoreCheckpointConfig # noqa: E402
from verl.workers.engine.fsdp.transformer_impl import FSDPEngineWithLMHead # noqa: E402

from trinity.common.config import Config # noqa: E402
from trinity.trainer.verl import monkey_patch # noqa: E402
from trinity.trainer.verl.config import ( # noqa: E402
_build_actor_config,
_build_critic_config,
_build_ref_config,
)


@pytest.mark.parametrize("strategy", ["fsdp", "fsdp2", "megatron"])
@pytest.mark.parametrize("role", ["actor", "ref", "critic"])
def test_worker_checkpoint_and_disabled_profiler_configs(strategy, role):
"""Build the exact nested dataclasses consumed by each worker role."""
config = Config()
sections = {
"actor": _build_actor_config(config, strategy, total_training_steps=10),
"ref": _build_ref_config(config, strategy),
"critic": _build_critic_config(config, strategy, use_critic=True, total_training_steps=10),
}
section = sections[role]
checkpoint = omega_conf_to_dataclass(OmegaConf.create(section["checkpoint"]))
expected_contents = ["model"] if role == "ref" else ["model", "optimizer", "extra"]
assert checkpoint.save_contents == expected_contents
assert checkpoint.load_contents == expected_contents
assert checkpoint.async_save is False
if strategy == "megatron":
assert isinstance(checkpoint, McoreCheckpointConfig)
assert checkpoint.mbridge_config == {
"distributed_filesystem": True,
"memory_efficient": True,
"strict": False,
}
else:
assert type(checkpoint) is CheckpointConfig
assert "mbridge_config" not in section["checkpoint"]

profiler = omega_conf_to_dataclass(OmegaConf.create(section["profiler"]))
assert profiler.enable is False
assert profiler.tool is None
# TrainingWorker performs this lookup even when profiling is disabled.
assert profiler.tool_config.get(profiler.tool, {}) == {}


def make_micro_batch():
"""Create two packed CPU sequences without loading a model or starting Ray."""
input_ids = torch.nested.nested_tensor_from_jagged(
torch.tensor([1, 2, 3, 4, 5]), offsets=torch.tensor([0, 3, 5])
)
position_ids = torch.nested.nested_tensor_from_jagged(
torch.tensor([0, 1, 2, 0, 1]), offsets=input_ids.offsets()
)
return TensorDict(
{
"input_ids": input_ids,
"position_ids": position_ids,
"temperature": torch.tensor([1.0, 2.0]),
},
batch_size=[2],
)


def test_packed_outputs_without_sequence_parallelism(monkeypatch):
"""SP=1 round-trips through verl's output preparation with zero padding."""
# Some test environments have flash-attn installed; keep this a CPU-only test.
monkeypatch.setenv("VERL_DISABLE_FLASH_ATTN_CE", "1")
batch = make_micro_batch()
engine = SimpleNamespace(use_ulysses_sp=False)
model_inputs, output_args = monkey_patch.prepare_model_inputs(engine, batch)
assert output_args["pad_size"] == 0
torch.testing.assert_close(
model_inputs["cu_seq_lens_q"], torch.tensor([0, 3, 5], dtype=torch.int32)
)
torch.testing.assert_close(
model_inputs["seq_idx"], torch.tensor([[0, 0, 0, 1, 1]], dtype=torch.int32)
)

engine._gather_and_unpad_packed = mock.Mock(side_effect=lambda value, pad_size: value)
logits = torch.arange(40, dtype=torch.float32).reshape(1, 5, 8) / 10
output = FSDPEngineWithLMHead.prepare_model_outputs(
engine, SimpleNamespace(logits=logits), output_args, batch, logits_processor_func=None
)
expected = torch.log_softmax(logits.squeeze(0) / torch.tensor([1, 1, 1, 2, 2])[:, None], dim=-1)
expected = expected.gather(1, torch.tensor([2, 3, 4, 5, 1])[:, None]).squeeze(1)
torch.testing.assert_close(output["log_probs"].values(), expected)
torch.testing.assert_close(output["log_probs"].offsets(), batch["input_ids"].offsets())
assert engine._gather_and_unpad_packed.call_args.args[1] == 0


@pytest.mark.parametrize("is_vlm", [False, True])
def test_sequence_parallel_padding_is_preserved(is_vlm):
"""Nonzero SP padding still extends sequence metadata for text and VLM paths."""
config = SimpleNamespace(vision_config={}) if is_vlm else SimpleNamespace()
engine = SimpleNamespace(
use_ulysses_sp=True, ulysses_sequence_parallel_size=2, module=SimpleNamespace(config=config)
)

def pad_inputs(inputs, position_ids_rmpad=None, **kwargs):
padding = 1
padded_inputs = torch.nn.functional.pad(
inputs, (0, padding), value=kwargs.get("pad_value", 0)
)
padded_positions = (
None
if position_ids_rmpad is None
else torch.nn.functional.pad(position_ids_rmpad, (0, padding))
)
return padded_inputs, padded_positions, padding

with mock.patch.object(monkey_patch, "ulysses_pad", side_effect=pad_inputs), mock.patch.object(
monkey_patch, "ulysses_pad_and_slice_inputs", side_effect=pad_inputs
):
model_inputs, output_args = monkey_patch.prepare_model_inputs(engine, make_micro_batch())

assert output_args["pad_size"] == 1
assert output_args["input_ids_rmpad_rolled"].shape == (6,)
torch.testing.assert_close(
model_inputs["cu_seq_lens_q"], torch.tensor([0, 3, 6], dtype=torch.int32)
)
torch.testing.assert_close(
model_inputs["seq_idx"], torch.tensor([[0, 0, 0, 1, 1, 1]], dtype=torch.int32)
)
22 changes: 14 additions & 8 deletions trinity/trainer/verl/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -311,7 +311,7 @@ def _build_actor_config(cfg: Config, strategy: str, total_training_steps: int) -
},
"router_replay": {"mode": "R3" if cfg.algorithm.enable_router_replay else "disabled"},
"profiler": _build_profiler_config(),
"checkpoint": _build_checkpoint_config(),
"checkpoint": _build_checkpoint_config(include_mbridge_config=is_megatron),
"optim": _build_optimizer_config(cfg.algorithm.optimizer, strategy, total_training_steps),
}

Expand Down Expand Up @@ -361,7 +361,9 @@ def _build_ref_config(cfg: Config, strategy: str) -> dict:
"use_prefix_grouper": False,
"profiler": _build_profiler_config(),
"router_replay": {"mode": "disabled"},
"checkpoint": _build_checkpoint_config(save_contents=["model"], load_contents=["model"]),
"checkpoint": _build_checkpoint_config(
save_contents=["model"], load_contents=["model"], include_mbridge_config=is_megatron
),
}

# Strategy-specific fields
Expand Down Expand Up @@ -469,7 +471,7 @@ def _build_critic_config(
"rollout_n": cfg.algorithm.repeat_times,
"profiler": _build_profiler_config(),
"optim": _build_critic_optimizer_config(strategy, total_training_steps),
"checkpoint": _build_checkpoint_config(),
"checkpoint": _build_checkpoint_config(include_mbridge_config=is_megatron),
}

# Strategy-specific fields
Expand Down Expand Up @@ -647,30 +649,34 @@ def _build_profiler_config() -> dict:
"all_ranks": False,
"ranks": [],
"save_path": "outputs/profile",
"tool_config": None,
"tool_config": {},
"global_tool_config": None,
}


def _build_checkpoint_config(
save_contents: Optional[List[str]] = None,
load_contents: Optional[List[str]] = None,
include_mbridge_config: bool = False,
) -> dict:
"""Build CheckpointConfig-compatible dict."""
if save_contents is None:
save_contents = ["model", "optimizer", "extra"]
if load_contents is None:
load_contents = ["model", "optimizer", "extra"]
return {
config = {
"save_contents": save_contents,
"load_contents": load_contents,
"async_save": False,
"mbridge_config": {
}
if include_mbridge_config:
# `mbridge_config` only exists on McoreCheckpointConfig (megatron backend)
config["mbridge_config"] = {
"distributed_filesystem": True,
"memory_efficient": True,
"strict": False,
},
}
}
return config


# ---------------------------------------------------------------------------
Expand Down
2 changes: 2 additions & 0 deletions trinity/trainer/verl/monkey_patch.py
Original file line number Diff line number Diff line change
Expand Up @@ -414,6 +414,8 @@ def prepare_model_inputs(self, micro_batch: TensorDict):
# for compute the log_prob
input_ids_rmpad_rolled = torch.roll(input_ids_rmpad, shifts=-1, dims=1) # (1, total_nnz)

# The output preparation reads pad_size even when sequence parallelism is disabled.
output_args["pad_size"] = 0
# pad and slice the inputs if sp > 1
if self.use_ulysses_sp:
is_vlm_model = hasattr(
Expand Down
Loading