Skip to content
Open
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
12 changes: 12 additions & 0 deletions docs/en/advanced/rl-kernel-linear-logp.md
Original file line number Diff line number Diff line change
Expand Up @@ -59,6 +59,8 @@ train/rl_kernel_linear_logp_fallback
train/rl_kernel_linear_logp_backend_descriptor_id
train/rl_kernel_linear_logp_contract_descriptor_id
train/rl_kernel_linear_logp_fallback_reason_descriptor_id
train/rl_kernel_linear_logp_fallback_reason_code_descriptor_id
train/rl_kernel_linear_logp_strict_failure
train/rl_kernel_linear_logp_call_count_total
train/rl_kernel_linear_logp_call_count_delta
train/rl_kernel_linear_logp_token_count_total
Expand Down Expand Up @@ -101,11 +103,21 @@ scheduling.
Native fallback is preserved. Unsupported cases record `fallback=1` and a
structured fallback reason, then materialize logits and compute selected
logprobs through the native vime/Megatron path unless strict mode is enabled.
Under `--rlk-fast strict`, the same unsupported cases emit a `strict-failure`
execution decision and raise instead of marking a native fallback.

Common fallback reasons include:

- the optional RL-Kernel package is unavailable;
- `rollout_temperature` is not `1.0`;
- strict mode is missing active response loss masks, or the masks are not
binary and response-length aligned;
- entropy was requested;
- CP redistribution is active;
- tensor-parallel metadata such as `tp_group`, `vocab_start_index`, or
`global_vocab_size` is incomplete;
- dtype/downcast metadata does not satisfy the strict contract;
- the selected op does not accept tensor-parallel metadata;
- the selected op does not return an autograd-connected tensor for train-time
full-gradient use;
- the model output layer or LM-head weight is unavailable.
184 changes: 184 additions & 0 deletions tests/test_rl_kernel_linear_logp_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,18 @@ def __call__(
return torch.gather(torch.log_softmax(logits, dim=-1), -1, target_ids.long().unsqueeze(-1)).squeeze(-1)


class _FakeDetachedLinearLogpOp(_FakeLinearLogpOp):
def __call__(
self,
hidden: torch.Tensor,
weight: torch.Tensor,
target_ids: torch.Tensor,
bias: torch.Tensor | None = None,
**kwargs,
) -> torch.Tensor:
return super().__call__(hidden, weight, target_ids, bias, **kwargs).detach()


def _drop_rl_engine_modules() -> None:
for name in list(sys.modules):
if name == "rl_engine" or name.startswith("rl_engine."):
Expand Down Expand Up @@ -254,6 +266,178 @@ def test_linear_logp_full_gradient_path_matches_materialized_logits(monkeypatch)
torch.testing.assert_close(actual_grads[2], bias_ref.grad, rtol=1e-6, atol=1e-6)


@pytest.mark.unit
def test_linear_logp_strict_fast_success_records_no_fallback(monkeypatch):
_install_fake_rl_engine(monkeypatch)
args = _make_args(rlk_fast="strict")
torch.manual_seed(12)
hidden = torch.randn(4, 3, requires_grad=True)
weight = torch.randn(6, 3, requires_grad=True)
target = torch.randint(0, 6, (4,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None)

actual = rlk_mod.maybe_compute_linear_logp(
hidden,
target,
context=context,
args=args,
with_entropy=False,
loss_masks=[torch.tensor([1, 1]), torch.tensor([1, 0])],
response_lengths=[2, 2],
)

torch.testing.assert_close(actual, _reference_logp(hidden, weight, target, None))
assert rlk_mod.get_rl_kernel_fallback_count("linear_logp") == 0
counters = rlk_mod.get_rl_kernel_runtime_counters()
assert counters["linear_logp_call_count"] == 1.0
assert counters["linear_logp_fallback_count"] == 0.0
metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["actual_backend"] == "_FakeLinearLogpOp"
assert metadata["fallback"] is False
assert metadata["fallback_reason"] is None
assert metadata["fallback_reason_code"] is None
assert metadata["strict_failure"] is False
assert rlk_mod.get_linear_logp_runtime_log_metrics(prefix="x/")["x/strict_failure"] == 0.0


@pytest.mark.unit
def test_linear_logp_auto_fallback_has_structured_temperature_reason(monkeypatch):
_install_fake_rl_engine(monkeypatch)
args = _make_args(rollout_temperature=0.7)
hidden = torch.randn(3, 4)
weight = torch.randn(6, 4)
target = torch.randint(0, 6, (3,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None)

actual = rlk_mod.maybe_compute_linear_logp(hidden, target, context=context, args=args, with_entropy=False)

assert actual is None
assert _FakeLinearLogpOp.calls == []
assert rlk_mod.get_rl_kernel_fallback_count("linear_logp") == 1
metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["actual_backend"] == "vime.native.linear_logp"
assert metadata["fallback"] is True
assert metadata["fallback_reason_code"] == "unsupported_temperature"
assert metadata["strict_failure"] is False


@pytest.mark.unit
def test_linear_logp_strict_fails_when_temperature_would_fallback(monkeypatch):
_install_fake_rl_engine(monkeypatch)
args = _make_args(rlk_fast="strict", rollout_temperature=0.7)
hidden = torch.randn(3, 4)
weight = torch.randn(6, 4)
target = torch.randint(0, 6, (3,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None)

with pytest.raises(RuntimeError, match="rollout_temperature=1.0"):
rlk_mod.maybe_compute_linear_logp(hidden, target, context=context, args=args, with_entropy=False)

assert _FakeLinearLogpOp.calls == []
assert rlk_mod.get_rl_kernel_fallback_count("linear_logp") == 0
metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["actual_backend"] is None
assert metadata["fallback"] is False
assert metadata["fallback_reason_code"] == "unsupported_temperature"
assert metadata["strict_failure"] is True


@pytest.mark.unit
def test_linear_logp_strict_requires_active_mask_metadata(monkeypatch):
_install_fake_rl_engine(monkeypatch)
args = _make_args(rlk_fast="strict")
hidden = torch.randn(2, 4)
weight = torch.randn(5, 4)
target = torch.randint(0, 5, (2,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None)

with pytest.raises(RuntimeError, match="active response loss masks"):
rlk_mod.maybe_compute_linear_logp(hidden, target, context=context, args=args, with_entropy=False)

metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["fallback_reason_code"] == "active_mask_missing"
assert metadata["strict_failure"] is True


@pytest.mark.unit
def test_linear_logp_strict_rejects_invalid_active_mask(monkeypatch):
_install_fake_rl_engine(monkeypatch)
args = _make_args(rlk_fast="strict")
hidden = torch.randn(2, 4)
weight = torch.randn(5, 4)
target = torch.randint(0, 5, (2,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None)

with pytest.raises(RuntimeError, match="0/1"):
rlk_mod.maybe_compute_linear_logp(
hidden,
target,
context=context,
args=args,
with_entropy=False,
loss_masks=[torch.tensor([1, 2])],
response_lengths=[2],
)

metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["fallback_reason_code"] == "active_mask_not_binary"
assert metadata["strict_failure"] is True


@pytest.mark.unit
def test_linear_logp_strict_requires_tensor_parallel_metadata(monkeypatch):
_install_fake_rl_engine(monkeypatch)
mpu.get_tensor_model_parallel_world_size.return_value = 2
mpu.get_tensor_model_parallel_group.return_value = None
args = _make_args(rlk_fast="strict")
hidden = torch.randn(2, 4)
weight = torch.randn(5, 4)
target = torch.randint(0, 10, (2,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None, global_vocab_size=10)

with pytest.raises(RuntimeError, match="tensor-parallel group metadata"):
rlk_mod.maybe_compute_linear_logp(
hidden,
target,
context=context,
args=args,
with_entropy=False,
loss_masks=[torch.tensor([1, 1])],
response_lengths=[2],
)

metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["fallback_reason_code"] == "tp_group_missing"
assert metadata["strict_failure"] is True


@pytest.mark.unit
def test_linear_logp_strict_rejects_missing_backward_saved_state(monkeypatch):
_install_fake_rl_engine(monkeypatch, op_factory=lambda: _FakeDetachedLinearLogpOp())
args = _make_args(rlk_fast="strict")
hidden = torch.randn(2, 4, requires_grad=True)
weight = torch.randn(5, 4, requires_grad=True)
target = torch.randint(0, 5, (2,))
context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=None, tp_group=None)

with pytest.raises(RuntimeError, match="autograd-connected"):
rlk_mod.maybe_compute_linear_logp(
hidden,
target,
context=context,
args=args,
with_entropy=False,
loss_masks=[torch.tensor([1, 1])],
response_lengths=[2],
)

assert len(_FakeLinearLogpOp.calls) == 1
metadata = rlk_mod.get_linear_logp_runtime_metadata()
assert metadata["fallback_reason_code"] == "backward_saved_state_missing"
assert metadata["fallback"] is False
assert metadata["strict_failure"] is True


@pytest.mark.unit
def test_linear_logp_matches_vime_response_slicing_from_hidden_states(monkeypatch):
_install_fake_rl_engine(monkeypatch)
Expand Down
5 changes: 5 additions & 0 deletions vime/backends/megatron_utils/loss.py
Original file line number Diff line number Diff line change
Expand Up @@ -543,6 +543,7 @@ def get_log_probs_and_entropy(
unconcat_tokens: list[torch.Tensor],
total_lengths: list[int],
response_lengths: list[int],
loss_masks: list[torch.Tensor] | None = None,
with_entropy: bool = False,
non_loss_data: bool = True,
top_p_token_ids: list[list[int]] | None = None,
Expand Down Expand Up @@ -593,6 +594,8 @@ def get_log_probs_and_entropy(
context=linear_logp_context,
args=args,
with_entropy=with_entropy,
loss_masks=loss_masks,
response_lengths=response_lengths,
)

if log_prob_full is None and linear_logp_context is not None:
Expand Down Expand Up @@ -1014,6 +1017,7 @@ def policy_loss_function(
unconcat_tokens=batch["unconcat_tokens"],
total_lengths=total_lengths,
response_lengths=response_lengths,
loss_masks=batch["loss_masks"],
with_entropy=need_entropy,
rl_kernel_linear_logp_context=rl_kernel_linear_logp_context,
**get_rollout_top_p_logprob_kwargs(args, batch),
Expand Down Expand Up @@ -1319,6 +1323,7 @@ def sft_loss_function(
unconcat_tokens=batch["unconcat_tokens"],
total_lengths=total_lengths,
response_lengths=response_lengths,
loss_masks=batch.get("loss_masks"),
with_entropy=False,
rl_kernel_linear_logp_context=rl_kernel_linear_logp_context,
)
Expand Down
4 changes: 3 additions & 1 deletion vime/backends/megatron_utils/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
from megatron.core.pipeline_parallel.utils import unwrap_model
except ImportError:
from megatron.core.utils import unwrap_model

from vime.utils import logging_utils
from vime.utils.memory_utils import clear_memory
from vime.utils.rl_kernel import is_rl_kernel_op_enabled
Expand All @@ -39,8 +40,8 @@
from .model_provider import get_model_provider_func
from .rl_kernel import (
get_linear_logp_context_from_model,
get_linear_logp_runtime_metadata,
get_linear_logp_runtime_log_metrics,
get_linear_logp_runtime_metadata,
get_rl_kernel_runtime_counter_delta,
get_rl_kernel_runtime_counters,
return_hidden_states_for_linear_logp,
Expand Down Expand Up @@ -482,6 +483,7 @@ def forward_step(
"unconcat_tokens": unconcat_tokens,
"total_lengths": total_lengths,
"response_lengths": response_lengths,
"loss_masks": batch["loss_masks"],
"with_entropy": args.use_rollout_entropy,
}
if use_rollout_top_p_replay:
Expand Down
Loading