From c21a49dd8782ad5f60b892868893af8d1ed5c988 Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Tue, 14 Jul 2026 20:28:50 +0800 Subject: [PATCH 1/3] [Core] Add RL-Kernel mode controls Co-authored-by: OpenAI Codex Signed-off-by: inaniloquentee <3051000145@qq.com> --- tests/test_megatron_argument_validation.py | 134 +++++++++++++++++++++ vime/utils/arguments.py | 121 +++++++++++++++++++ 2 files changed, 255 insertions(+) diff --git a/tests/test_megatron_argument_validation.py b/tests/test_megatron_argument_validation.py index 0e96b140..72706074 100644 --- a/tests/test_megatron_argument_validation.py +++ b/tests/test_megatron_argument_validation.py @@ -1,3 +1,4 @@ +import argparse import importlib.util import sys import types @@ -176,6 +177,134 @@ def test_update_weight_disk_dir_required_for_disk_transport(monkeypatch): module.vime_validate_args(args) +@pytest.mark.unit +def test_add_rl_kernel_arguments_registers_controls(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + parser = argparse.ArgumentParser(add_help=False) + + module.add_rl_kernel_arguments(parser) + + flags = {flag for action in parser._actions for flag in action.option_strings} + assert "--rlk-fast" in flags + assert "--rlk-consistency" in flags + assert "--enable-rl-kernel" in flags + assert "--rl-kernel-strict" in flags + assert "--rl-kernel-ops" in flags + + +@pytest.mark.unit +def test_rl_kernel_help_text_describes_orthogonal_controls(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + parser = argparse.ArgumentParser(add_help=False) + + module.add_rl_kernel_arguments(parser) + help_text = " ".join(parser.format_help().split()) + + assert "independently from consistency auditing" in help_text + assert "independently from fast-path acceleration" in help_text + + +def make_rlk_args(**overrides): + values = dict( + rlk_fast=None, + rlk_consistency=None, + enable_rl_kernel=False, + rl_kernel_strict=False, + rl_kernel_ops=(), + ) + values.update(overrides) + return types.SimpleNamespace(**values) + + +@pytest.mark.unit +def test_resolve_rlk_mode_defaults_to_native_off(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + monkeypatch.delenv("VIME_RLK_FAST", raising=False) + monkeypatch.delenv("VIME_RLK_CONSISTENCY", raising=False) + args = make_rlk_args() + + config = module.resolve_rlk_mode_config(args) + + assert config.fast == "off" + assert config.consistency == "off" + assert config.ops == () + assert args.rlk_fast == "off" + assert args.rlk_consistency == "off" + assert args.rlk_mode_config == config + + +@pytest.mark.unit +def test_resolve_rlk_mode_uses_env_aliases(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + monkeypatch.setenv("VIME_RLK_FAST", "strict") + monkeypatch.setenv("VIME_RLK_CONSISTENCY", "audit") + args = make_rlk_args() + + config = module.resolve_rlk_mode_config(args) + + assert config.fast == "strict" + assert config.consistency == "audit" + + +@pytest.mark.unit +def test_resolve_rlk_mode_cli_overrides_env_and_legacy_flags(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + monkeypatch.setenv("VIME_RLK_FAST", "strict") + monkeypatch.setenv("VIME_RLK_CONSISTENCY", "audit") + args = make_rlk_args( + rlk_fast="off", + rlk_consistency="strict", + enable_rl_kernel=True, + rl_kernel_strict=True, + rl_kernel_ops=("linear_logp",), + ) + + config = module.resolve_rlk_mode_config(args) + + assert config.fast == "off" + assert config.consistency == "strict" + assert config.ops == ("linear_logp",) + + +@pytest.mark.unit +def test_resolve_rlk_mode_legacy_enable_maps_to_auto(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + monkeypatch.delenv("VIME_RLK_FAST", raising=False) + args = make_rlk_args(enable_rl_kernel=True) + + assert module.resolve_rlk_mode_config(args).fast == "auto" + + +@pytest.mark.unit +def test_resolve_rlk_mode_legacy_strict_maps_to_strict(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + monkeypatch.delenv("VIME_RLK_FAST", raising=False) + args = make_rlk_args(enable_rl_kernel=True, rl_kernel_strict=True) + + assert module.resolve_rlk_mode_config(args).fast == "strict" + + +@pytest.mark.unit +def test_resolve_rlk_mode_rejects_invalid_env(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + monkeypatch.setenv("VIME_RLK_FAST", "turbo") + args = make_rlk_args() + + with pytest.raises(ValueError, match="VIME_RLK_FAST"): + module.resolve_rlk_mode_config(args) + + +@pytest.mark.unit +def test_rl_kernel_ops_parse_comma_allowlist(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + parser = argparse.ArgumentParser(add_help=False) + module.add_rl_kernel_arguments(parser) + + args = parser.parse_args(["--rl-kernel-ops", "linear_logp, logp,,attention"]) + + assert args.rl_kernel_ops == ("linear_logp", "logp", "attention") + + def make_vime_validate_args(**overrides): values = dict( eval_config=None, @@ -250,6 +379,11 @@ def make_vime_validate_args(**overrides): rollout_max_context_len=None, rollout_max_prompt_len=None, train_backend="megatron", + rlk_fast=None, + rlk_consistency=None, + enable_rl_kernel=False, + rl_kernel_strict=False, + rl_kernel_ops=(), release_train=False, keep_old_actor=False, only_train_params_name_list=None, diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 25fd9cc9..0329249c 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -3,6 +3,7 @@ import json import logging import os +from dataclasses import dataclass from typing import Any import yaml @@ -16,6 +17,16 @@ logger = logging.getLogger(__name__) +RLK_FAST_CHOICES = ("off", "auto", "strict") +RLK_CONSISTENCY_CHOICES = ("off", "audit", "strict") + + +@dataclass(frozen=True) +class RlkModeConfig: + fast: str + consistency: str + ops: tuple[str, ...] + def reset_arg(parser, name, **kwargs): """ @@ -33,6 +44,113 @@ def reset_arg(parser, name, **kwargs): parser.add_argument(name, **kwargs) +def _parse_rl_kernel_ops(value): + if value is None: + return () + if isinstance(value, tuple): + return value + if isinstance(value, list): + return tuple(value) + return tuple(op.strip() for op in value.split(",") if op.strip()) + + +def _get_env_choice(name, choices): + raw_value = os.environ.get(name) + if raw_value is None or raw_value == "": + return None + + value = raw_value.strip().lower() + if value not in choices: + allowed = ", ".join(choices) + raise ValueError(f"{name} must be one of {{{allowed}}}, got {raw_value!r}.") + return value + + +def _validate_rlk_choice(name, value, choices): + if value not in choices: + allowed = ", ".join(choices) + raise ValueError(f"--{name.replace('_', '-')} must be one of {{{allowed}}}, got {value!r}.") + return value + + +def resolve_rlk_mode_config(args): + """Resolve RL-Kernel public controls into one config object.""" + env_fast = ( + None if getattr(args, "rlk_fast", None) is not None else _get_env_choice("VIME_RLK_FAST", RLK_FAST_CHOICES) + ) + env_consistency = ( + None + if getattr(args, "rlk_consistency", None) is not None + else _get_env_choice("VIME_RLK_CONSISTENCY", RLK_CONSISTENCY_CHOICES) + ) + + legacy_fast = None + if getattr(args, "rl_kernel_strict", False): + legacy_fast = "strict" + elif getattr(args, "enable_rl_kernel", False): + legacy_fast = "auto" + + fast = _validate_rlk_choice( + "rlk_fast", getattr(args, "rlk_fast", None) or env_fast or legacy_fast or "off", RLK_FAST_CHOICES + ) + consistency = _validate_rlk_choice( + "rlk_consistency", + getattr(args, "rlk_consistency", None) or env_consistency or "off", + RLK_CONSISTENCY_CHOICES, + ) + ops = _parse_rl_kernel_ops(getattr(args, "rl_kernel_ops", ())) + + config = RlkModeConfig(fast=fast, consistency=consistency, ops=ops) + args.rlk_fast = fast + args.rlk_consistency = consistency + args.rl_kernel_ops = ops + args.rlk_mode_config = config + return config + + +def add_rl_kernel_arguments(parser): + parser.add_argument( + "--rlk-fast", + choices=RLK_FAST_CHOICES, + default=None, + help=( + "Select RL-Kernel fast-path behavior independently from consistency auditing: " + "'off' keeps native vime execution, 'auto' uses eligible RL-Kernel backends with fallback, " + "and 'strict' requires enabled RL-Kernel backends. Defaults to VIME_RLK_FAST or off; " + "--enable-rl-kernel maps to auto and --rl-kernel-strict maps to strict when not explicitly set." + ), + ) + parser.add_argument( + "--rlk-consistency", + choices=RLK_CONSISTENCY_CHOICES, + default=None, + help=( + "Select rollout-training consistency diagnostics independently from fast-path acceleration: " + "'off' disables diagnostics, 'audit' reports diagnostics without changing execution, " + "and 'strict' requires contract checks. Defaults to VIME_RLK_CONSISTENCY or off." + ), + ) + parser.add_argument( + "--enable-rl-kernel", + action="store_true", + default=False, + help="Compatibility alias: selects --rlk-fast auto unless --rlk-fast or VIME_RLK_FAST is set.", + ) + parser.add_argument( + "--rl-kernel-strict", + action="store_true", + default=False, + help="Compatibility alias: selects --rlk-fast strict unless --rlk-fast or VIME_RLK_FAST is set.", + ) + parser.add_argument( + "--rl-kernel-ops", + type=_parse_rl_kernel_ops, + default=(), + help="Comma-separated RL-Kernel operator allowlist, for example linear_logp,logp.", + ) + return parser + + def get_vime_extra_args_provider(add_custom_arguments=None): def add_vime_arguments(parser): # Ray @@ -1516,6 +1634,7 @@ def add_ci_arguments(parser): parser = add_cluster_arguments(parser) parser = add_train_arguments(parser) + parser = add_rl_kernel_arguments(parser) parser = add_rollout_arguments(parser) parser = add_fault_tolerance_arguments(parser) parser = add_data_arguments(parser) @@ -1747,6 +1866,7 @@ def _validate_update_weight_args(args) -> None: def vime_validate_args(args): + resolve_rlk_mode_config(args) args.eval_datasets = _resolve_eval_datasets(args) if args.kl_coef != 0 or args.use_kl_loss: @@ -2009,6 +2129,7 @@ def vime_validate_args(args): if hasattr(args, k): logger.info(f"Warning: Argument {k} is already set to {getattr(args, k)}, will override with {v}.") setattr(args, k, v) + resolve_rlk_mode_config(args) if args.eval_max_context_len is None: logger.info( From f77ad27183b036ba4f16f0f454932c449f123b17 Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Sun, 19 Jul 2026 14:57:57 +0800 Subject: [PATCH 2/3] Address RL-Kernel mode parser review Co-authored-by: OpenAI Codex Signed-off-by: inaniloquentee <3051000145@qq.com> --- tests/test_megatron_argument_validation.py | 24 ++++++++++++++++++++++ vime/utils/arguments.py | 4 ++++ 2 files changed, 28 insertions(+) diff --git a/tests/test_megatron_argument_validation.py b/tests/test_megatron_argument_validation.py index 72706074..9cd0d266 100644 --- a/tests/test_megatron_argument_validation.py +++ b/tests/test_megatron_argument_validation.py @@ -192,6 +192,30 @@ def test_add_rl_kernel_arguments_registers_controls(monkeypatch): assert "--rl-kernel-ops" in flags +@pytest.mark.unit +def test_standard_vime_parser_registers_rl_kernel_controls(monkeypatch): + module = load_vime_arguments_module(monkeypatch) + parser = argparse.ArgumentParser(add_help=False) + + module.get_vime_extra_args_provider()(parser) + args = parser.parse_args( + [ + "--rlk-fast", + "auto", + "--rlk-consistency", + "audit", + "--rl-kernel-ops", + "linear_logp,logp", + "--rollout-batch-size", + "8", + ] + ) + + assert args.rlk_fast == "auto" + assert args.rlk_consistency == "audit" + assert args.rl_kernel_ops == ("linear_logp", "logp") + + @pytest.mark.unit def test_rl_kernel_help_text_describes_orthogonal_controls(monkeypatch): module = load_vime_arguments_module(monkeypatch) diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 0329249c..a218f765 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -101,6 +101,8 @@ def resolve_rlk_mode_config(args): ops = _parse_rl_kernel_ops(getattr(args, "rl_kernel_ops", ())) config = RlkModeConfig(fast=fast, consistency=consistency, ops=ops) + # Compatibility aliases for existing callers; new code should prefer the + # immutable args.rlk_mode_config so resolved mode state has one owner. args.rlk_fast = fast args.rlk_consistency = consistency args.rl_kernel_ops = ops @@ -1634,6 +1636,8 @@ def add_ci_arguments(parser): parser = add_cluster_arguments(parser) parser = add_train_arguments(parser) + # Production parser path: RL-Kernel controls are framework-level vime + # flags, not CI-only switches. parser = add_rl_kernel_arguments(parser) parser = add_rollout_arguments(parser) parser = add_fault_tolerance_arguments(parser) From af7e995defefda521f1658a2e91f9a48d263387f Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Mon, 20 Jul 2026 22:58:21 +0800 Subject: [PATCH 3/3] Add RL-Kernel operator adapter boundary Signed-off-by: inaniloquentee <3051000145@qq.com> --- .../en/advanced/rl-kernel-operator-adapter.md | 24 + docs/en/index.rst | 1 + tests/test_rl_kernel_operator_adapter.py | 394 ++++++++++ vime/backends/rl_kernel_utils/__init__.py | 53 ++ vime/backends/rl_kernel_utils/adapter.py | 685 ++++++++++++++++++ 5 files changed, 1157 insertions(+) create mode 100644 docs/en/advanced/rl-kernel-operator-adapter.md create mode 100644 tests/test_rl_kernel_operator_adapter.py create mode 100644 vime/backends/rl_kernel_utils/__init__.py create mode 100644 vime/backends/rl_kernel_utils/adapter.py diff --git a/docs/en/advanced/rl-kernel-operator-adapter.md b/docs/en/advanced/rl-kernel-operator-adapter.md new file mode 100644 index 00000000..b65d9cc2 --- /dev/null +++ b/docs/en/advanced/rl-kernel-operator-adapter.md @@ -0,0 +1,24 @@ +# RL-Kernel Operator Adapter + +`vime.backends.rl_kernel_utils` is the vime-owned boundary for optional RL-Kernel-backed operator calls. Production vime code should depend on `RlkOperatorAdapter`, `NoOpRlkOperatorAdapter`, or the input dataclasses from this package, not on RL-Kernel internals. + +The boundary has three jobs: + +- carry resolved vime policy into the adapter (`fast`, `consistency`, and `enabled_ops`); +- map borrowed vime runtime tensors and metadata into operator-call inputs without taking ownership of rollout, training, Ray, Megatron, or artifact lifecycles; +- isolate the only optional RL-Kernel registry touchpoint, `rl_engine.kernels.registry.kernel_registry.get_op(...)`. + +When RL-Kernel modes are disabled, `build_rlk_operator_adapter(...)` returns `NoOpRlkOperatorAdapter`. The no-op adapter reports disabled decisions and lets the native vime path continue unchanged. + +For tests, use `MockRlkOperatorAdapter`. It implements the same protocol without importing RL-Kernel, CUDA, Triton, or distributed runtime packages. + +The current Phase 1 hooks are: + +- `capability(op_name)`; +- `contract(op_name, runtime=..., metadata=...)`; +- `selected_logprobs(SelectedLogprobInputs)`; +- `reference_logprobs(ReferenceScoreInputs)`; +- `linear_logp(LinearLogpInputs)`; +- `provenance()`. + +Future `linear_logp` integration work should extend this adapter package rather than adding direct `rl_engine` imports to Megatron, rollout, or training modules. diff --git a/docs/en/index.rst b/docs/en/index.rst index b345724f..7825277b 100644 --- a/docs/en/index.rst +++ b/docs/en/index.rst @@ -66,6 +66,7 @@ Start by Use Case advanced/delta-weight-sync.md advanced/vllm-config.md advanced/megatron-config.md + advanced/rl-kernel-operator-adapter.md advanced/arch-support-beyond-megatron.md .. toctree:: diff --git a/tests/test_rl_kernel_operator_adapter.py b/tests/test_rl_kernel_operator_adapter.py new file mode 100644 index 00000000..809043f2 --- /dev/null +++ b/tests/test_rl_kernel_operator_adapter.py @@ -0,0 +1,394 @@ +import ast +import importlib +import sys +import types +import warnings +from pathlib import Path + +import pytest + + +def _fresh_adapter_module(): + sys.modules.pop("vime.backends.rl_kernel_utils", None) + sys.modules.pop("vime.backends.rl_kernel_utils.adapter", None) + return importlib.import_module("vime.backends.rl_kernel_utils.adapter") + + +def _drop_rl_engine_modules(): + for name in list(sys.modules): + if name == "rl_engine" or name.startswith("rl_engine."): + sys.modules.pop(name, None) + + +@pytest.mark.unit +def test_adapter_import_does_not_import_rl_engine(): + _drop_rl_engine_modules() + + _fresh_adapter_module() + + assert not any(name == "rl_engine" or name.startswith("rl_engine.") for name in sys.modules) + + +@pytest.mark.unit +def test_policy_context_is_built_from_resolved_vime_args(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + args = types.SimpleNamespace( + rlk_mode_config=types.SimpleNamespace(fast="auto", consistency="audit", ops=("linear_logp", "logp")), + ) + + policy = module.rlk_policy_context_from_args(args) + + assert policy.fast == "auto" + assert policy.consistency == "audit" + assert policy.enabled_ops == ("linear_logp", "logp") + assert policy.operator_enabled("linear_logp") + assert policy.operator_enabled("reference_logp") + + +@pytest.mark.unit +def test_policy_context_from_args_tolerates_unresolved_none_values(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + args = types.SimpleNamespace(rlk_fast=None, rlk_consistency=None, rl_kernel_ops=None) + + policy = module.rlk_policy_context_from_args(args) + + assert policy.fast == "off" + assert policy.consistency == "off" + assert policy.enabled_ops == () + + +@pytest.mark.unit +def test_noop_adapter_preserves_native_behavior_when_disabled(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + policy = module.RlkPolicyContext(fast="off", consistency="off", enabled_ops=("linear_logp",)) + + adapter = module.build_rlk_operator_adapter(policy) + result = adapter.linear_logp( + module.LinearLogpInputs(hidden="hidden", lm_head_weight="weight", target_ids=[1, 2, 3]) + ) + + assert isinstance(adapter, module.NoOpRlkOperatorAdapter) + assert result.value is None + assert result.decision.path == "disabled" + assert result.decision.token_count == 3 + assert adapter.telemetry.fallback_counts == {"linear_logp": 1} + + +@pytest.mark.unit +def test_consistency_audit_without_fast_path_keeps_operator_calls_native_only(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + policy = module.RlkPolicyContext(fast="off", consistency="audit", enabled_ops=("linear_logp", "logp")) + + adapter = module.build_rlk_operator_adapter(policy) + capability = adapter.capability("linear_logp") + result = adapter.linear_logp(module.LinearLogpInputs(hidden="hidden", lm_head_weight="weight", target_ids=[1, 2])) + + assert isinstance(adapter, module.NoOpRlkOperatorAdapter) + assert not policy.operator_enabled("linear_logp") + assert not capability.available + assert result.value is None + assert result.decision.path == "disabled" + + +@pytest.mark.unit +def test_mock_adapter_captures_linear_logp_inputs_and_runtime_metadata(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + loss_masks = object() + batch = { + "total_lengths": [5, 7], + "response_lengths": [2, 3], + "loss_masks": loss_masks, + "rollout_log_probs": ["old-logp"], + "rollout_top_p_token_ids": ["ids"], + "rollout_top_p_token_offsets": ["offsets"], + } + policy = module.RlkPolicyContext(fast="auto", consistency="audit", enabled_ops=("linear_logp",)) + adapter = module.MockRlkOperatorAdapter( + policy, + available_ops=("linear_logp",), + handlers={"linear_logp": lambda payload: (payload.hidden, payload.runtime.total_lengths, payload.metadata)}, + ) + + inputs = module.linear_logp_inputs_from_vime( + hidden="hidden-ref", + lm_head_weight="weight-ref", + target_ids=[4, 5], + batch=batch, + tp_group="tp-group", + vocab_start_index=8, + global_vocab_size=16, + metadata={"weight_version": "actor@1"}, + ) + result = adapter.linear_logp(inputs) + + assert inputs.runtime.total_lengths == (5, 7) + assert inputs.runtime.response_lengths == (2, 3) + assert inputs.runtime.loss_masks is loss_masks + assert inputs.tp_group == "tp-group" + assert inputs.vocab_start_index == 8 + assert inputs.global_vocab_size == 16 + assert result.value == ("hidden-ref", (5, 7), {"weight_version": "actor@1"}) + assert result.decision.path == "mock" + assert adapter.telemetry.call_counts == {"linear_logp": 1} + assert adapter.telemetry.token_counts == {"linear_logp": 2} + + contract = adapter.contract("linear_logp", runtime=inputs.runtime, metadata={"source": "unit"}) + assert contract.op_name == "linear_logp" + assert contract.runtime is inputs.runtime + assert contract.metadata == {"source": "unit"} + assert contract.policy["fast"] == "auto" + + +@pytest.mark.unit +def test_mock_adapter_reports_unsupported_operator_without_rl_kernel(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + policy = module.RlkPolicyContext(fast="auto", consistency="off", enabled_ops=("linear_logp",)) + adapter = module.MockRlkOperatorAdapter(policy, available_ops=("logp",)) + + result = adapter.linear_logp(module.LinearLogpInputs(hidden="hidden", lm_head_weight="weight", target_ids=[1, 2])) + + assert result.value is None + assert result.decision.path == "unsupported" + assert result.decision.token_count == 2 + assert "unsupported" in result.decision.reason + assert adapter.telemetry.fallback_counts == {"linear_logp": 1} + + +@pytest.mark.unit +def test_mock_reference_logprobs_reuses_logp_handler_when_reference_handler_is_absent(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + policy = module.RlkPolicyContext(fast="auto", consistency="audit", enabled_ops=("logp",)) + adapter = module.MockRlkOperatorAdapter( + policy, + available_ops=("logp",), + handlers={"logp": lambda payload: ("logp", payload.logits, payload.target_ids)}, + ) + + result = adapter.reference_logprobs(module.ReferenceScoreInputs(logits="ref-logits", target_ids=[3])) + + assert result.value == ("logp", "ref-logits", [3]) + assert result.decision.path == "mock" + assert adapter.telemetry.call_counts == {"reference_logp": 1} + + +@pytest.mark.unit +def test_registry_adapter_defers_rl_engine_import_until_capability_query(monkeypatch): + module = _fresh_adapter_module() + _drop_rl_engine_modules() + + adapter = module.build_rlk_operator_adapter( + module.RlkPolicyContext(fast="auto", consistency="off", enabled_ops=("linear_logp",)) + ) + + assert isinstance(adapter, module.RlkRegistryOperatorAdapter) + assert not any(name == "rl_engine" or name.startswith("rl_engine.") for name in sys.modules) + + class FakeLinearLogpOp: + def __init__(self): + self.calls = [] + + def __call__(self, *args, **kwargs): + self.calls.append((args, kwargs)) + return "rlk-output" + + fake_op = FakeLinearLogpOp() + + class FakeRegistry: + def __init__(self): + self.requested = [] + + def get_op(self, op_name): + self.requested.append(op_name) + assert op_name == "linear_logp" + return fake_op + + fake_registry = FakeRegistry() + rl_engine_mod = types.ModuleType("rl_engine") + kernels_mod = types.ModuleType("rl_engine.kernels") + registry_mod = types.ModuleType("rl_engine.kernels.registry") + registry_mod.kernel_registry = fake_registry + monkeypatch.setitem(sys.modules, "rl_engine", rl_engine_mod) + monkeypatch.setitem(sys.modules, "rl_engine.kernels", kernels_mod) + monkeypatch.setitem(sys.modules, "rl_engine.kernels.registry", registry_mod) + + result = adapter.linear_logp( + module.LinearLogpInputs( + hidden="hidden", + lm_head_weight="weight", + target_ids=[9], + bias="bias", + tp_group="tp", + vocab_start_index=4, + global_vocab_size=12, + ) + ) + + assert fake_registry.requested == ["linear_logp"] + assert result.value == "rlk-output" + assert result.decision.path == "fast" + assert result.decision.backend == "FakeLinearLogpOp" + assert fake_op.calls == [ + ( + ("hidden", "weight", [9], "bias"), + {"tp_group": "tp", "vocab_start_index": 4, "global_vocab_size": 12}, + ) + ] + + +@pytest.mark.unit +def test_registry_logprob_hooks_accept_int_masks(monkeypatch): + torch = pytest.importorskip("torch") + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + + class FakeLogpOp: + def __call__(self, logits, target_ids): + return torch.tensor([-0.1, -0.2, -0.3], dtype=torch.float32) + + class FakeRegistry: + def get_op(self, op_name): + assert op_name == "logp" + return FakeLogpOp() + + registry_mod = types.ModuleType("rl_engine.kernels.registry") + registry_mod.kernel_registry = FakeRegistry() + monkeypatch.setitem(sys.modules, "rl_engine", types.ModuleType("rl_engine")) + monkeypatch.setitem(sys.modules, "rl_engine.kernels", types.ModuleType("rl_engine.kernels")) + monkeypatch.setitem(sys.modules, "rl_engine.kernels.registry", registry_mod) + + adapter = module.RlkRegistryOperatorAdapter( + module.RlkPolicyContext(fast="auto", consistency="audit", enabled_ops=("logp",)) + ) + inputs = module.SelectedLogprobInputs( + logits="unused", + target_ids=torch.tensor([1, 2, 3]), + mask=torch.tensor([1, 0, 1], dtype=torch.int32), + ) + + result = adapter.selected_logprobs(inputs) + + torch.testing.assert_close(result.value, torch.tensor([-0.1, 0.0, -0.3], dtype=torch.float32)) + + +@pytest.mark.unit +def test_registry_adapter_strict_mode_raises_when_enabled_op_is_unavailable(monkeypatch): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + + class MissingRegistry: + def get_op(self, op_name): + raise RuntimeError(f"{op_name} missing") + + registry_mod = types.ModuleType("rl_engine.kernels.registry") + registry_mod.kernel_registry = MissingRegistry() + monkeypatch.setitem(sys.modules, "rl_engine", types.ModuleType("rl_engine")) + monkeypatch.setitem(sys.modules, "rl_engine.kernels", types.ModuleType("rl_engine.kernels")) + monkeypatch.setitem(sys.modules, "rl_engine.kernels.registry", registry_mod) + + adapter = module.RlkRegistryOperatorAdapter( + module.RlkPolicyContext(fast="strict", consistency="off", enabled_ops=("linear_logp",)) + ) + + with pytest.raises(module.RlkOperatorUnavailable, match="linear_logp"): + adapter.linear_logp(module.LinearLogpInputs(hidden="h", lm_head_weight="w", target_ids=[1])) + + +@pytest.mark.unit +def test_registry_adapter_fallback_records_token_count_when_op_is_unavailable(monkeypatch): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + + class MissingRegistry: + def get_op(self, op_name): + raise RuntimeError(f"{op_name} missing") + + registry_mod = types.ModuleType("rl_engine.kernels.registry") + registry_mod.kernel_registry = MissingRegistry() + monkeypatch.setitem(sys.modules, "rl_engine", types.ModuleType("rl_engine")) + monkeypatch.setitem(sys.modules, "rl_engine.kernels", types.ModuleType("rl_engine.kernels")) + monkeypatch.setitem(sys.modules, "rl_engine.kernels.registry", registry_mod) + + adapter = module.RlkRegistryOperatorAdapter( + module.RlkPolicyContext(fast="auto", consistency="off", enabled_ops=("linear_logp",)) + ) + result = adapter.linear_logp(module.LinearLogpInputs(hidden="h", lm_head_weight="w", target_ids=[1, 2, 3])) + + assert result.value is None + assert result.decision.path == "fallback" + assert result.decision.token_count == 3 + assert adapter.telemetry.fallback_counts == {"linear_logp": 1} + assert adapter.telemetry.token_counts == {"linear_logp": 3} + + +@pytest.mark.unit +def test_selected_and_reference_logprob_hooks_share_boundary_contract(): + module = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + policy = module.RlkPolicyContext(fast="auto", consistency="audit", enabled_ops=("logp",)) + adapter = module.MockRlkOperatorAdapter( + policy, + available_ops=("logp",), + handlers={ + "logp": lambda payload: ("selected", payload.logits, payload.target_ids), + "reference_logp": lambda payload: ("reference", payload.logits, payload.target_ids), + }, + ) + + selected = adapter.selected_logprobs(module.SelectedLogprobInputs(logits="logits", target_ids=[1, 2])) + reference = adapter.reference_logprobs(module.ReferenceScoreInputs(logits="ref-logits", target_ids=[3])) + + assert selected.value == ("selected", "logits", [1, 2]) + assert reference.value == ("reference", "ref-logits", [3]) + assert adapter.telemetry.call_counts == {"logp": 1, "reference_logp": 1} + + +@pytest.mark.unit +def test_production_vime_code_does_not_import_rl_engine_outside_adapter_boundary(): + repo = Path(__file__).resolve().parents[1] + allowed = Path("vime/backends/rl_kernel_utils/adapter.py") + bad_imports = [] + + for package_root in (repo / "vime", repo / "vime_plugins"): + for path in package_root.rglob("*.py"): + rel = path.relative_to(repo).as_posix() + if rel == allowed.as_posix() or "__pycache__" in path.parts: + continue + with warnings.catch_warnings(): + warnings.filterwarnings("ignore", category=DeprecationWarning, message="invalid escape sequence.*") + tree = ast.parse(path.read_text(encoding="utf-8"), filename=str(path)) + for node in ast.walk(tree): + if isinstance(node, ast.Import): + for alias in node.names: + if alias.name == "rl_engine" or alias.name.startswith("rl_engine."): + bad_imports.append((rel, node.lineno, alias.name)) + elif isinstance(node, ast.ImportFrom): + module_name = node.module or "" + if module_name == "rl_engine" or module_name.startswith("rl_engine."): + bad_imports.append((rel, node.lineno, module_name)) + + assert bad_imports == [] + + +@pytest.mark.unit +def test_adapter_boundary_only_references_rl_kernel_public_registry(): + repo = Path(__file__).resolve().parents[1] + source = (repo / "vime" / "backends" / "rl_kernel_utils" / "adapter.py").read_text(encoding="utf-8") + + assert '"rl_engine.kernels.registry"' in source + forbidden_fragments = [ + "PairedRunner", + "RuntimeTools", + "cross_config", + "child_process", + "artifact_layout", + "attempt_artifact", + ] + assert [fragment for fragment in forbidden_fragments if fragment in source] == [] + + +@pytest.mark.unit +def test_rl_kernel_adapter_boundary_is_documented(): + repo = Path(__file__).resolve().parents[1] + doc = repo / "docs" / "en" / "advanced" / "rl-kernel-operator-adapter.md" + text = doc.read_text(encoding="utf-8") + + assert "RlkOperatorAdapter" in text + assert "vime.backends.rl_kernel_utils" in text + assert "rl_engine.kernels.registry" in text + assert "contract(op_name" in text diff --git a/vime/backends/rl_kernel_utils/__init__.py b/vime/backends/rl_kernel_utils/__init__.py new file mode 100644 index 00000000..8c1e5daa --- /dev/null +++ b/vime/backends/rl_kernel_utils/__init__.py @@ -0,0 +1,53 @@ +"""RL-Kernel adapter boundary owned by vime.""" + +from vime.backends.rl_kernel_utils.adapter import ( + RLK_ALL_OPERATORS, + RLK_OP_LINEAR_LOGP, + RLK_OP_REFERENCE_LOGPROBS, + RLK_OP_SELECTED_LOGPROBS, + LinearLogpInputs, + MockRlkOperatorAdapter, + NoOpRlkOperatorAdapter, + ReferenceScoreInputs, + RlkCapability, + RlkOperatorAdapter, + RlkOperatorContract, + RlkOperatorDecision, + RlkOperatorResult, + RlkOperatorTelemetry, + RlkOperatorUnavailable, + RlkPolicyContext, + RlkRegistryOperatorAdapter, + RlkRuntimeBatchMetadata, + SelectedLogprobInputs, + build_rlk_operator_adapter, + linear_logp_inputs_from_vime, + rlk_policy_context_from_args, + runtime_batch_metadata_from_vime_batch, +) + +__all__ = [ + "LinearLogpInputs", + "MockRlkOperatorAdapter", + "NoOpRlkOperatorAdapter", + "ReferenceScoreInputs", + "RLK_ALL_OPERATORS", + "RLK_OP_LINEAR_LOGP", + "RLK_OP_REFERENCE_LOGPROBS", + "RLK_OP_SELECTED_LOGPROBS", + "RlkCapability", + "RlkOperatorContract", + "RlkOperatorAdapter", + "RlkOperatorDecision", + "RlkOperatorResult", + "RlkOperatorTelemetry", + "RlkOperatorUnavailable", + "RlkPolicyContext", + "RlkRegistryOperatorAdapter", + "RlkRuntimeBatchMetadata", + "SelectedLogprobInputs", + "build_rlk_operator_adapter", + "linear_logp_inputs_from_vime", + "rlk_policy_context_from_args", + "runtime_batch_metadata_from_vime_batch", +] diff --git a/vime/backends/rl_kernel_utils/adapter.py b/vime/backends/rl_kernel_utils/adapter.py new file mode 100644 index 00000000..333cede4 --- /dev/null +++ b/vime/backends/rl_kernel_utils/adapter.py @@ -0,0 +1,685 @@ +"""vime-owned boundary for optional RL-Kernel operator calls. + +This module is the only production vime surface that should know how to ask +RL-Kernel for an operator implementation. Importing it is deliberately cheap: +it does not import ``rl_engine`` and it does not initialize CUDA, Triton, Ray, +Megatron, or any cross-config runner machinery. +""" + +from __future__ import annotations + +import importlib +import time +from collections.abc import Callable, Mapping +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import Any, Protocol, runtime_checkable + +RLK_OP_LINEAR_LOGP = "linear_logp" +RLK_OP_SELECTED_LOGPROBS = "logp" +RLK_OP_REFERENCE_LOGPROBS = "reference_logp" +RLK_ALL_OPERATORS = "*" + +_FAST_CHOICES = {"off", "auto", "strict"} +_CONSISTENCY_CHOICES = {"off", "audit", "strict"} + + +def _immutable_mapping(value: Mapping[str, Any] | None) -> Mapping[str, Any]: + return MappingProxyType(dict(value or {})) + + +def _normalize_ops(value: Any) -> tuple[str, ...]: + if value is None: + return () + if isinstance(value, str): + return tuple(op.strip() for op in value.split(",") if op.strip()) + return tuple(str(op).strip() for op in value if str(op).strip()) + + +@dataclass(frozen=True) +class RlkPolicyContext: + """Resolved vime policy passed into an RL-Kernel adapter. + + The adapter consumes this object; it never parses CLI flags or environment + variables itself. ``enabled_ops`` is an allowlist. Use ``("*",)`` only for + tests or explicitly opted-in future configs that want every known hook. + """ + + fast: str = "off" + consistency: str = "off" + enabled_ops: tuple[str, ...] = () + metadata: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + fast = str(self.fast).strip().lower() + consistency = str(self.consistency).strip().lower() + if fast not in _FAST_CHOICES: + raise ValueError(f"fast must be one of {sorted(_FAST_CHOICES)}, got {self.fast!r}.") + if consistency not in _CONSISTENCY_CHOICES: + raise ValueError(f"consistency must be one of {sorted(_CONSISTENCY_CHOICES)}, got {self.consistency!r}.") + object.__setattr__(self, "fast", fast) + object.__setattr__(self, "consistency", consistency) + object.__setattr__(self, "enabled_ops", _normalize_ops(self.enabled_ops)) + object.__setattr__(self, "metadata", _immutable_mapping(self.metadata)) + + @property + def native_only(self) -> bool: + return self.fast == "off" and self.consistency == "off" + + @property + def strict_fast(self) -> bool: + return self.fast == "strict" + + def operator_enabled(self, op_name: str) -> bool: + if self.fast == "off": + return False + if RLK_ALL_OPERATORS in self.enabled_ops: + return True + if op_name in self.enabled_ops: + return True + # Reference scoring reuses selected-logprob implementations unless a + # future backend advertises a distinct reference scorer. + return op_name == RLK_OP_REFERENCE_LOGPROBS and RLK_OP_SELECTED_LOGPROBS in self.enabled_ops + + +def rlk_policy_context_from_args(args: Any) -> RlkPolicyContext: + """Build a policy context from already-resolved vime args. + + ``vime.utils.arguments.resolve_rlk_mode_config`` owns parsing and alias + resolution. This helper only reads the resolved attributes so callers can + pass a compact, immutable policy into adapter construction. + """ + + config = getattr(args, "rlk_mode_config", None) + if config is not None: + fast = config.fast + consistency = config.consistency + enabled_ops = config.ops + else: + fast = getattr(args, "rlk_fast", None) or "off" + consistency = getattr(args, "rlk_consistency", None) or "off" + enabled_ops = getattr(args, "rl_kernel_ops", None) or () + return RlkPolicyContext( + fast=fast, + consistency=consistency, + enabled_ops=enabled_ops, + metadata={"source": "vime_args"}, + ) + + +@dataclass(frozen=True) +class RlkRuntimeBatchMetadata: + """Borrowed rollout/training metadata for operator calls. + + Tensor/list fields are stored by reference. The adapter does not own rollout + samples, data iterators, training lifecycle, or cross-process artifacts. + """ + + total_lengths: tuple[int, ...] = () + response_lengths: tuple[int, ...] = () + loss_masks: Any | None = None + rollout_log_probs: Any | None = None + rollout_top_p_token_ids: Any | None = None + rollout_top_p_token_offsets: Any | None = None + metadata: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "total_lengths", tuple(int(x) for x in self.total_lengths)) + object.__setattr__(self, "response_lengths", tuple(int(x) for x in self.response_lengths)) + object.__setattr__(self, "metadata", _immutable_mapping(self.metadata)) + + +def runtime_batch_metadata_from_vime_batch(batch: Mapping[str, Any] | None) -> RlkRuntimeBatchMetadata: + """Map a vime rollout/training batch to borrowed RL-Kernel call metadata.""" + + if not batch: + return RlkRuntimeBatchMetadata() + return RlkRuntimeBatchMetadata( + total_lengths=tuple(batch.get("total_lengths", ())), + response_lengths=tuple(batch.get("response_lengths", ())), + loss_masks=batch.get("loss_masks"), + rollout_log_probs=batch.get("rollout_log_probs"), + rollout_top_p_token_ids=batch.get("rollout_top_p_token_ids"), + rollout_top_p_token_offsets=batch.get("rollout_top_p_token_offsets"), + metadata=batch.get("metadata") or {}, + ) + + +@dataclass(frozen=True) +class SelectedLogprobInputs: + logits: Any + target_ids: Any + mask: Any | None = None + temperature: float = 1.0 + runtime: RlkRuntimeBatchMetadata = field(default_factory=RlkRuntimeBatchMetadata) + metadata: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "metadata", _immutable_mapping(self.metadata)) + + +@dataclass(frozen=True) +class ReferenceScoreInputs: + logits: Any + target_ids: Any + mask: Any | None = None + runtime: RlkRuntimeBatchMetadata = field(default_factory=RlkRuntimeBatchMetadata) + metadata: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "metadata", _immutable_mapping(self.metadata)) + + +@dataclass(frozen=True) +class LinearLogpInputs: + hidden: Any + lm_head_weight: Any + target_ids: Any + bias: Any | None = None + tp_group: Any | None = None + vocab_start_index: int = 0 + global_vocab_size: int | None = None + runtime: RlkRuntimeBatchMetadata = field(default_factory=RlkRuntimeBatchMetadata) + metadata: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "vocab_start_index", int(self.vocab_start_index)) + if self.global_vocab_size is not None: + object.__setattr__(self, "global_vocab_size", int(self.global_vocab_size)) + object.__setattr__(self, "metadata", _immutable_mapping(self.metadata)) + + +def linear_logp_inputs_from_vime( + *, + hidden: Any, + lm_head_weight: Any, + target_ids: Any, + bias: Any | None = None, + batch: Mapping[str, Any] | None = None, + tp_group: Any | None = None, + vocab_start_index: int = 0, + global_vocab_size: int | None = None, + metadata: Mapping[str, Any] | None = None, +) -> LinearLogpInputs: + """Create a ``linear_logp`` call payload from vime-owned runtime objects.""" + + return LinearLogpInputs( + hidden=hidden, + lm_head_weight=lm_head_weight, + target_ids=target_ids, + bias=bias, + tp_group=tp_group, + vocab_start_index=vocab_start_index, + global_vocab_size=global_vocab_size, + runtime=runtime_batch_metadata_from_vime_batch(batch), + metadata=metadata or {}, + ) + + +@dataclass(frozen=True) +class RlkCapability: + op_name: str + available: bool + backend: str = "none" + reason: str | None = None + provenance: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "provenance", _immutable_mapping(self.provenance)) + + +@dataclass(frozen=True) +class RlkOperatorContract: + op_name: str + policy: Mapping[str, Any] + runtime: RlkRuntimeBatchMetadata | None = None + metadata: Mapping[str, Any] = field(default_factory=dict) + provenance: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "policy", _immutable_mapping(self.policy)) + object.__setattr__(self, "metadata", _immutable_mapping(self.metadata)) + object.__setattr__(self, "provenance", _immutable_mapping(self.provenance)) + + +@dataclass(frozen=True) +class RlkOperatorDecision: + op_name: str + path: str + backend: str = "none" + reason: str | None = None + elapsed_s: float = 0.0 + token_count: int | None = None + provenance: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "provenance", _immutable_mapping(self.provenance)) + + +@dataclass(frozen=True) +class RlkOperatorResult: + value: Any + decision: RlkOperatorDecision + + +@dataclass +class RlkOperatorTelemetry: + capability_queries: dict[str, int] = field(default_factory=dict) + call_counts: dict[str, int] = field(default_factory=dict) + fallback_counts: dict[str, int] = field(default_factory=dict) + token_counts: dict[str, int] = field(default_factory=dict) + dispatch_elapsed_s: dict[str, float] = field(default_factory=dict) + last_decision: RlkOperatorDecision | None = None + + def record_capability_query(self, op_name: str) -> None: + self.capability_queries[op_name] = self.capability_queries.get(op_name, 0) + 1 + + def record_decision(self, decision: RlkOperatorDecision) -> None: + op_name = decision.op_name + self.last_decision = decision + if decision.path in {"fast", "reference", "mock"}: + self.call_counts[op_name] = self.call_counts.get(op_name, 0) + 1 + if decision.path in {"fallback", "disabled", "unsupported"}: + self.fallback_counts[op_name] = self.fallback_counts.get(op_name, 0) + 1 + if decision.token_count is not None: + self.token_counts[op_name] = self.token_counts.get(op_name, 0) + int(decision.token_count) + self.dispatch_elapsed_s[op_name] = self.dispatch_elapsed_s.get(op_name, 0.0) + float(decision.elapsed_s) + + def snapshot(self) -> dict[str, Any]: + return { + "capability_queries": dict(self.capability_queries), + "call_counts": dict(self.call_counts), + "fallback_counts": dict(self.fallback_counts), + "token_counts": dict(self.token_counts), + "dispatch_elapsed_s": dict(self.dispatch_elapsed_s), + "last_decision": None if self.last_decision is None else self.last_decision, + } + + +class RlkOperatorUnavailable(RuntimeError): + """Raised when strict vime policy requires an unavailable RL-Kernel op.""" + + +@runtime_checkable +class RlkOperatorAdapter(Protocol): + policy: RlkPolicyContext + telemetry: RlkOperatorTelemetry + + def capability(self, op_name: str, *, metadata: Mapping[str, Any] | None = None) -> RlkCapability: ... + + def contract( + self, + op_name: str, + *, + runtime: RlkRuntimeBatchMetadata | None = None, + metadata: Mapping[str, Any] | None = None, + ) -> RlkOperatorContract: ... + + def selected_logprobs(self, inputs: SelectedLogprobInputs) -> RlkOperatorResult: ... + + def reference_logprobs(self, inputs: ReferenceScoreInputs) -> RlkOperatorResult: ... + + def linear_logp(self, inputs: LinearLogpInputs) -> RlkOperatorResult: ... + + def provenance(self) -> Mapping[str, Any]: ... + + +class NoOpRlkOperatorAdapter: + """Disabled adapter that keeps native vime execution unchanged.""" + + name = "noop" + + def __init__(self, policy: RlkPolicyContext | None = None) -> None: + self.policy = policy or RlkPolicyContext() + self.telemetry = RlkOperatorTelemetry() + + def capability(self, op_name: str, *, metadata: Mapping[str, Any] | None = None) -> RlkCapability: + self.telemetry.record_capability_query(op_name) + reason = "RL-Kernel operator path is disabled by vime policy." + return RlkCapability(op_name=op_name, available=False, reason=reason, provenance=metadata or {}) + + def contract( + self, + op_name: str, + *, + runtime: RlkRuntimeBatchMetadata | None = None, + metadata: Mapping[str, Any] | None = None, + ) -> RlkOperatorContract: + return RlkOperatorContract( + op_name=op_name, + policy=_policy_payload(self.policy), + runtime=runtime, + metadata=metadata or {}, + provenance=self.provenance(), + ) + + def _disabled_result(self, op_name: str, token_count: int | None = None) -> RlkOperatorResult: + decision = RlkOperatorDecision( + op_name=op_name, + path="disabled", + reason="RL-Kernel operator path is disabled by vime policy.", + token_count=token_count, + provenance=self.provenance(), + ) + self.telemetry.record_decision(decision) + return RlkOperatorResult(value=None, decision=decision) + + def selected_logprobs(self, inputs: SelectedLogprobInputs) -> RlkOperatorResult: + return self._disabled_result(RLK_OP_SELECTED_LOGPROBS, _num_tokens_from_targets(inputs.target_ids)) + + def reference_logprobs(self, inputs: ReferenceScoreInputs) -> RlkOperatorResult: + return self._disabled_result(RLK_OP_REFERENCE_LOGPROBS, _num_tokens_from_targets(inputs.target_ids)) + + def linear_logp(self, inputs: LinearLogpInputs) -> RlkOperatorResult: + return self._disabled_result(RLK_OP_LINEAR_LOGP, _num_tokens_from_targets(inputs.target_ids)) + + def provenance(self) -> Mapping[str, Any]: + return _immutable_mapping( + { + "adapter": self.name, + "fast": self.policy.fast, + "consistency": self.policy.consistency, + "enabled_ops": self.policy.enabled_ops, + } + ) + + +class RlkRegistryOperatorAdapter: + """Optional adapter backed by RL-Kernel's public kernel registry.""" + + name = "registry" + + def __init__(self, policy: RlkPolicyContext) -> None: + self.policy = policy + self.telemetry = RlkOperatorTelemetry() + self._ops: dict[str, Any] = {} + + def _registry(self) -> Any: + module = importlib.import_module("rl_engine.kernels.registry") + return module.kernel_registry + + def _get_op(self, op_name: str) -> Any: + if op_name not in self._ops: + registry_op_name = RLK_OP_SELECTED_LOGPROBS if op_name == RLK_OP_REFERENCE_LOGPROBS else op_name + self._ops[op_name] = self._registry().get_op(registry_op_name) + return self._ops[op_name] + + def capability(self, op_name: str, *, metadata: Mapping[str, Any] | None = None) -> RlkCapability: + self.telemetry.record_capability_query(op_name) + if not self.policy.operator_enabled(op_name): + return RlkCapability( + op_name=op_name, + available=False, + reason=f"{op_name!r} is not enabled by vime RL-Kernel policy.", + provenance=metadata or {}, + ) + try: + op = self._get_op(op_name) + except Exception as exc: + return RlkCapability( + op_name=op_name, + available=False, + reason=f"RL-Kernel registry could not provide {op_name!r}: {exc}", + provenance=metadata or {}, + ) + return RlkCapability( + op_name=op_name, + available=True, + backend=type(op).__name__, + provenance={"adapter": self.name, **dict(metadata or {})}, + ) + + def contract( + self, + op_name: str, + *, + runtime: RlkRuntimeBatchMetadata | None = None, + metadata: Mapping[str, Any] | None = None, + ) -> RlkOperatorContract: + return RlkOperatorContract( + op_name=op_name, + policy=_policy_payload(self.policy), + runtime=runtime, + metadata=metadata or {}, + provenance=self.provenance(), + ) + + def _unsupported_result_or_raise( + self, + op_name: str, + capability: RlkCapability, + *, + token_count: int | None = None, + ) -> RlkOperatorResult: + path = "unsupported" if self.policy.strict_fast else "fallback" + decision = RlkOperatorDecision( + op_name=op_name, + path=path, + backend=capability.backend, + reason=capability.reason, + token_count=token_count, + provenance=self.provenance(), + ) + self.telemetry.record_decision(decision) + if self.policy.strict_fast: + raise RlkOperatorUnavailable(capability.reason or f"RL-Kernel op {op_name!r} is unavailable.") + return RlkOperatorResult(value=None, decision=decision) + + def _call(self, op_name: str, token_count: int | None, fn: Callable[[Any], Any]) -> RlkOperatorResult: + capability = self.capability(op_name) + if not capability.available: + return self._unsupported_result_or_raise(op_name, capability, token_count=token_count) + op = self._get_op(op_name) + start = time.perf_counter() + value = fn(op) + elapsed_s = time.perf_counter() - start + decision = RlkOperatorDecision( + op_name=op_name, + path="fast", + backend=type(op).__name__, + elapsed_s=elapsed_s, + token_count=token_count, + provenance=self.provenance(), + ) + self.telemetry.record_decision(decision) + return RlkOperatorResult(value=value, decision=decision) + + def selected_logprobs(self, inputs: SelectedLogprobInputs) -> RlkOperatorResult: + if inputs.temperature != 1.0: + capability = RlkCapability( + op_name=RLK_OP_SELECTED_LOGPROBS, + available=False, + reason="selected_logprobs adapter path does not own temperature replay yet.", + ) + return self._unsupported_result_or_raise( + RLK_OP_SELECTED_LOGPROBS, + capability, + token_count=_num_tokens_from_targets(inputs.target_ids), + ) + + def invoke(op: Any) -> Any: + value = op(inputs.logits, inputs.target_ids) + return _mask_inactive_values(value, inputs.mask) + + return self._call(RLK_OP_SELECTED_LOGPROBS, _num_tokens_from_targets(inputs.target_ids), invoke) + + def reference_logprobs(self, inputs: ReferenceScoreInputs) -> RlkOperatorResult: + def invoke(op: Any) -> Any: + value = op(inputs.logits, inputs.target_ids) + return _mask_inactive_values(value, inputs.mask) + + return self._call(RLK_OP_REFERENCE_LOGPROBS, _num_tokens_from_targets(inputs.target_ids), invoke) + + def linear_logp(self, inputs: LinearLogpInputs) -> RlkOperatorResult: + def invoke(op: Any) -> Any: + return op( + inputs.hidden, + inputs.lm_head_weight, + inputs.target_ids, + inputs.bias, + tp_group=inputs.tp_group, + vocab_start_index=inputs.vocab_start_index, + global_vocab_size=inputs.global_vocab_size, + ) + + return self._call(RLK_OP_LINEAR_LOGP, _num_tokens_from_targets(inputs.target_ids), invoke) + + def provenance(self) -> Mapping[str, Any]: + return _immutable_mapping( + { + "adapter": self.name, + "fast": self.policy.fast, + "consistency": self.policy.consistency, + "enabled_ops": self.policy.enabled_ops, + "boundary": "vime.backends.rl_kernel_utils", + } + ) + + +class MockRlkOperatorAdapter: + """Test adapter that satisfies ``RlkOperatorAdapter`` without RL-Kernel.""" + + name = "mock" + + def __init__( + self, + policy: RlkPolicyContext | None = None, + *, + available_ops: tuple[str, ...] = (RLK_ALL_OPERATORS,), + handlers: Mapping[str, Callable[[Any], Any]] | None = None, + ) -> None: + self.policy = policy or RlkPolicyContext(fast="auto", enabled_ops=available_ops) + self.available_ops = set(available_ops) + self.handlers = dict(handlers or {}) + self.telemetry = RlkOperatorTelemetry() + + def capability(self, op_name: str, *, metadata: Mapping[str, Any] | None = None) -> RlkCapability: + self.telemetry.record_capability_query(op_name) + enabled = self.policy.operator_enabled(op_name) + available = enabled and self._mock_op_available(op_name) + return RlkCapability( + op_name=op_name, + available=available, + backend=self.name if available else "none", + reason=None if available else f"{op_name!r} is unsupported by the mock adapter.", + provenance=metadata or {}, + ) + + def _mock_op_available(self, op_name: str) -> bool: + return ( + RLK_ALL_OPERATORS in self.available_ops + or op_name in self.available_ops + or (op_name == RLK_OP_REFERENCE_LOGPROBS and RLK_OP_SELECTED_LOGPROBS in self.available_ops) + ) + + def contract( + self, + op_name: str, + *, + runtime: RlkRuntimeBatchMetadata | None = None, + metadata: Mapping[str, Any] | None = None, + ) -> RlkOperatorContract: + return RlkOperatorContract( + op_name=op_name, + policy=_policy_payload(self.policy), + runtime=runtime, + metadata=metadata or {}, + provenance=self.provenance(), + ) + + def _call(self, op_name: str, payload: Any, token_count: int | None) -> RlkOperatorResult: + capability = self.capability(op_name) + if not capability.available: + decision = RlkOperatorDecision( + op_name=op_name, + path="unsupported", + backend="none", + reason=capability.reason, + token_count=token_count, + provenance=self.provenance(), + ) + self.telemetry.record_decision(decision) + return RlkOperatorResult(value=None, decision=decision) + handler = self.handlers.get(op_name) or self.handlers.get(_registry_op_name(op_name), lambda value: value) + value = handler(payload) + decision = RlkOperatorDecision( + op_name=op_name, + path="mock", + backend=self.name, + token_count=token_count, + provenance=self.provenance(), + ) + self.telemetry.record_decision(decision) + return RlkOperatorResult(value=value, decision=decision) + + def selected_logprobs(self, inputs: SelectedLogprobInputs) -> RlkOperatorResult: + return self._call(RLK_OP_SELECTED_LOGPROBS, inputs, _num_tokens_from_targets(inputs.target_ids)) + + def reference_logprobs(self, inputs: ReferenceScoreInputs) -> RlkOperatorResult: + return self._call(RLK_OP_REFERENCE_LOGPROBS, inputs, _num_tokens_from_targets(inputs.target_ids)) + + def linear_logp(self, inputs: LinearLogpInputs) -> RlkOperatorResult: + return self._call(RLK_OP_LINEAR_LOGP, inputs, _num_tokens_from_targets(inputs.target_ids)) + + def provenance(self) -> Mapping[str, Any]: + return _immutable_mapping( + { + "adapter": self.name, + "fast": self.policy.fast, + "consistency": self.policy.consistency, + "enabled_ops": self.policy.enabled_ops, + "available_ops": tuple(sorted(self.available_ops)), + } + ) + + +def build_rlk_operator_adapter( + policy: RlkPolicyContext | None = None, + *, + args: Any | None = None, + backend: str = "auto", +) -> RlkOperatorAdapter: + """Small construction point for selecting a concrete RL-Kernel adapter.""" + + if policy is None: + policy = rlk_policy_context_from_args(args) if args is not None else RlkPolicyContext() + if policy.fast == "off" or not policy.enabled_ops: + return NoOpRlkOperatorAdapter(policy) + + normalized_backend = backend.strip().lower().replace("-", "_") + if normalized_backend in {"auto", "registry", "rl_kernel", "rlk"}: + return RlkRegistryOperatorAdapter(policy) + if normalized_backend in {"none", "noop", "no_op", "disabled"}: + return NoOpRlkOperatorAdapter(policy) + raise ValueError(f"Unknown RL-Kernel operator adapter backend: {backend!r}.") + + +def _num_tokens_from_targets(target_ids: Any) -> int | None: + if target_ids is None: + return None + if hasattr(target_ids, "numel"): + return int(target_ids.numel()) + try: + return len(target_ids) + except TypeError: + return None + + +def _policy_payload(policy: RlkPolicyContext) -> dict[str, Any]: + return { + "fast": policy.fast, + "consistency": policy.consistency, + "enabled_ops": policy.enabled_ops, + "metadata": dict(policy.metadata), + } + + +def _registry_op_name(op_name: str) -> str: + return RLK_OP_SELECTED_LOGPROBS if op_name == RLK_OP_REFERENCE_LOGPROBS else op_name + + +def _mask_inactive_values(value: Any, mask: Any | None) -> Any: + if mask is None: + return value + bool_mask = mask.bool() if hasattr(mask, "bool") else mask + return value.masked_fill(~bool_mask, 0.0)