From c326551c5a68425c42c12d07ef3e132347c827f9 Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Tue, 14 Jul 2026 21:47:19 +0800 Subject: [PATCH 1/5] [Core] Add RL-Kernel execution decision records Co-authored-by: OpenAI Codex Signed-off-by: inaniloquentee <3051000145@qq.com> --- tests/utils/test_rl_kernel_execution.py | 419 ++++++++++++++++ vime/backends/rl_kernel_utils/__init__.py | 27 ++ vime/backends/rl_kernel_utils/execution.py | 528 +++++++++++++++++++++ 3 files changed, 974 insertions(+) create mode 100644 tests/utils/test_rl_kernel_execution.py create mode 100644 vime/backends/rl_kernel_utils/__init__.py create mode 100644 vime/backends/rl_kernel_utils/execution.py diff --git a/tests/utils/test_rl_kernel_execution.py b/tests/utils/test_rl_kernel_execution.py new file mode 100644 index 00000000..b332aa70 --- /dev/null +++ b/tests/utils/test_rl_kernel_execution.py @@ -0,0 +1,419 @@ +import json +import logging + +import pytest + +from vime.backends.rl_kernel_utils import ( + BackendCapability, + FallbackReason, + LogprobContractMetadata, + NumericContract, + RlKernelCapabilities, + build_logprob_contract_decision, + emit_execution_decision, + query_rl_kernel_capabilities, + select_execution_decision, +) +from vime.backends.rl_kernel_utils.execution import RLK_DECISION_EVENT + + +def _contract(contract_id: str = "rlk.linear_logp.fp32") -> NumericContract: + return NumericContract( + contract_id=contract_id, + accumulation_dtype="fp32", + reduction_order="vocab-shard-local-then-global", + downcast_point="after-selected-logprob", + sharding_rule="tp-vocab", + merge_semantics="selected-token", + quantization_policy="none", + backend_id="rlk.linear_logp.fast", + tolerance_by_dtype={"bf16": {"source": "rl-kernel", "dtype": "bf16"}}, + ) + + +def _backend( + backend_id: str = "rlk.linear_logp.fast", + *, + operator: str = "linear_logp", + implementation_kind: str = "optimized", + strict_fast_eligible: bool = True, + dtypes: tuple[str, ...] = ("bf16",), + contract_id: str = "rlk.linear_logp.fp32", +) -> BackendCapability: + return BackendCapability( + operator=operator, + backend_id=backend_id, + implementation_kind=implementation_kind, + dtypes=dtypes, + hardware_targets=("cuda",), + autograd_modes=("full-gradient",), + parallel_modes=("tp",), + deterministic=True, + batch_invariant=True, + strict_fast_eligible=strict_fast_eligible, + runtime_fingerprint="runtime-1", + build_fingerprint="build-1", + config_lifecycle="process-start", + numeric_contract=_contract(contract_id), + ) + + +def _caps(*backends: BackendCapability) -> RlKernelCapabilities: + return RlKernelCapabilities(available=True, backends=tuple(backends)) + + +@pytest.mark.unit +def test_off_off_selects_native_without_fallback(): + decision = select_execution_decision(operator="linear_logp", stage="train_logprob", dtype="bf16") + + assert decision.decision == "native" + assert decision.actual_backend == "vime.native" + assert decision.fallback is False + assert decision.contract_id == "vime.native.linear_logp" + + +@pytest.mark.unit +@pytest.mark.parametrize("consistency_mode", ["audit", "strict"]) +def test_fast_off_consistency_modes_are_audit_only(consistency_mode): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="off", + requested_consistency=consistency_mode, + ) + + assert decision.decision == "audit-only" + assert decision.actual_backend == "vime.native" + assert decision.fallback is False + + +@pytest.mark.unit +def test_auto_missing_capabilities_falls_back_with_structured_reason(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + requested_consistency="off", + ) + + assert decision.decision == "fallback-native" + assert decision.fallback is True + assert decision.actual_backend == "vime.native" + assert decision.fallback_reason == FallbackReason( + code="capability_data_missing", + message="RL-Kernel capability data is unavailable.", + ) + + +@pytest.mark.unit +def test_strict_missing_capabilities_returns_strict_failure(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="strict", + requested_consistency="off", + ) + + assert decision.decision == "strict-failure" + assert decision.actual_backend is None + assert decision.fallback is False + assert decision.fallback_reason.code == "capability_data_missing" + + +@pytest.mark.unit +def test_auto_selects_matching_optimized_backend(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + capabilities=_caps(_backend()), + requested_backend="rlk.linear_logp.fast", + dtype="bf16", + parallel_context={"tp": 2}, + ) + + assert decision.decision == "optimized" + assert decision.actual_backend == "rlk.linear_logp.fast" + assert decision.capability_backend_id == "rlk.linear_logp.fast" + assert decision.contract_id == "rlk.linear_logp.fp32" + assert decision.parallel_context == {"tp": 2} + + +@pytest.mark.unit +def test_auto_strict_consistency_selects_strict_fast_backend(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + requested_consistency="strict", + capabilities=_caps(_backend(strict_fast_eligible=True)), + dtype="bf16", + ) + + assert decision.decision == "strict-fast" + assert decision.actual_backend == "rlk.linear_logp.fast" + assert decision.strict_eligible is True + + +@pytest.mark.unit +def test_auto_strict_consistency_uses_reference_when_no_strict_fast_backend(): + opportunistic = _backend(strict_fast_eligible=False) + reference = _backend( + "rlk.linear_logp.reference", + implementation_kind="reference", + strict_fast_eligible=False, + contract_id="rlk.linear_logp.reference.fp32", + ) + + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + requested_consistency="strict", + capabilities=_caps(opportunistic, reference), + dtype="bf16", + ) + + assert decision.decision == "strict-reference" + assert decision.actual_backend == "rlk.linear_logp.reference" + assert decision.strict_eligible is True + + +@pytest.mark.unit +def test_strict_fast_strict_consistency_does_not_silently_use_reference(): + reference = _backend( + "rlk.linear_logp.reference", + implementation_kind="reference", + strict_fast_eligible=False, + contract_id="rlk.linear_logp.reference.fp32", + ) + + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="strict", + requested_consistency="strict", + capabilities=_caps(reference), + dtype="bf16", + ) + + assert decision.decision == "strict-failure" + assert decision.actual_backend is None + assert decision.fallback_reason.code == "strict_backend_unavailable" + + +@pytest.mark.unit +def test_requested_backend_mismatch_falls_back_in_auto_mode(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + capabilities=_caps(_backend("rlk.other")), + requested_backend="rlk.missing", + dtype="bf16", + ) + + assert decision.decision == "fallback-native" + assert decision.fallback_reason.code == "backend_unavailable" + assert decision.fallback_reason.details["requested_backend"] == "rlk.missing" + + +@pytest.mark.unit +def test_dtype_mismatch_returns_structured_backend_unavailable_reason(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + capabilities=_caps(_backend(dtypes=("fp16",))), + dtype="bf16", + ) + + assert decision.decision == "fallback-native" + assert decision.fallback_reason.code == "backend_unavailable" + assert decision.fallback_reason.details["dtype"] == "bf16" + + +@pytest.mark.unit +def test_query_capabilities_without_provider_is_explicitly_unavailable(): + result = query_rl_kernel_capabilities() + + assert result.capabilities.available is False + assert result.fallback_reason.code == "capability_provider_missing" + + +@pytest.mark.unit +def test_query_capabilities_normalizes_dict_provider(): + result = query_rl_kernel_capabilities( + lambda: { + "available": True, + "runtime_fingerprint": "runtime-1", + "backends": [ + { + "operator": "linear_logp", + "backend_id": "rlk.linear_logp.fast", + "implementation_kind": "optimized", + "dtypes": ["bf16"], + "strict_fast_eligible": True, + "numeric_contract": { + "contract_id": "rlk.linear_logp.fp32", + "tolerance_by_dtype": {"bf16": {"source": "provider"}}, + }, + } + ], + } + ) + + assert result.capabilities.available is True + assert result.capabilities.runtime_fingerprint == "runtime-1" + assert result.capabilities.backends[0].backend_id == "rlk.linear_logp.fast" + assert result.capabilities.backends[0].numeric_contract.tolerance_for_dtype("bf16") == {"source": "provider"} + + +@pytest.mark.unit +def test_query_capabilities_provider_failure_is_structured(): + def provider(): + raise RuntimeError("boom") + + result = query_rl_kernel_capabilities(provider) + + assert result.capabilities.available is False + assert result.fallback_reason.code == "capability_query_failed" + assert "boom" in result.fallback_reason.details["error"] + + +@pytest.mark.unit +def test_tolerance_lookup_uses_contract_metadata_without_hardcoded_thresholds(): + capabilities = _caps(_backend()) + + assert capabilities.tolerance_for_contract("rlk.linear_logp.fp32", "bf16") == { + "source": "rl-kernel", + "dtype": "bf16", + } + + +@pytest.mark.unit +def test_tolerance_lookup_rejects_unknown_contract(): + capabilities = _caps(_backend()) + + with pytest.raises(KeyError, match="unknown RL-Kernel numeric contract"): + capabilities.tolerance_for_contract("missing", "bf16") + + +@pytest.mark.unit +def test_logprob_contract_match_returns_reportable_contract_identity(): + decision = build_logprob_contract_decision( + rollout_contract=LogprobContractMetadata( + source="rollout", + contract_id="rlk.linear_logp.fp32", + dtype="bf16", + backend_id="rlk.linear_logp.fast", + ), + recomputed_contract=LogprobContractMetadata( + source="train", + contract_id="rlk.linear_logp.fp32", + dtype="bf16", + backend_id="rlk.linear_logp.fast", + ), + strict=True, + ) + + assert decision.decision == "contract-match" + assert decision.contract_id == "rlk.linear_logp.fp32" + assert decision.strict_eligible is True + assert decision.details == { + "rollout_contract_id": "rlk.linear_logp.fp32", + "recomputed_contract_id": "rlk.linear_logp.fp32", + } + + +@pytest.mark.unit +@pytest.mark.parametrize("strict,expected_decision", [(False, "audit-warning"), (True, "strict-failure")]) +def test_logprob_contract_mismatch_is_structured(strict, expected_decision): + decision = build_logprob_contract_decision( + rollout_contract=LogprobContractMetadata(source="rollout", contract_id="rollout.contract"), + recomputed_contract=LogprobContractMetadata(source="train", contract_id="train.contract"), + strict=strict, + ) + + assert decision.decision == expected_decision + assert decision.fallback_reason.code == "contract_id_mismatch" + assert decision.fallback_reason.details == { + "rollout_contract_id": "rollout.contract", + "recomputed_contract_id": "train.contract", + } + + +@pytest.mark.unit +def test_logprob_contract_missing_fails_closed_in_strict_mode(): + decision = build_logprob_contract_decision( + rollout_contract=LogprobContractMetadata(source="rollout", contract_id=None), + recomputed_contract=LogprobContractMetadata(source="train", contract_id="train.contract"), + strict=True, + ) + + assert decision.decision == "strict-failure" + assert decision.fallback_reason.code == "contract_id_missing" + + +@pytest.mark.unit +def test_execution_decision_log_record_is_stable_and_json_serializable(): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + capabilities=_caps(_backend()), + dtype="bf16", + ) + + record = decision.to_log_record() + encoded = json.dumps(record, sort_keys=True) + + assert json.loads(encoded)["event"] == RLK_DECISION_EVENT + assert sorted(record) == [ + "actual_backend", + "capability_backend_id", + "contract_id", + "decision", + "details", + "dtype", + "event", + "fallback", + "fallback_reason", + "operator", + "parallel_context", + "requested_backend", + "requested_mode", + "stage", + "strict_eligible", + ] + + +@pytest.mark.unit +def test_emit_execution_decision_logs_json_payload(caplog): + decision = select_execution_decision( + operator="linear_logp", + stage="train_logprob", + requested_fast="auto", + capabilities=_caps(_backend()), + dtype="bf16", + ) + + with caplog.at_level(logging.INFO, logger="vime.backends.rl_kernel_utils.execution"): + emit_execution_decision(decision) + + assert RLK_DECISION_EVENT in caplog.text + payload = caplog.records[0].message.split(RLK_DECISION_EVENT, 1)[1].strip() + assert json.loads(payload)["actual_backend"] == "rlk.linear_logp.fast" + + +@pytest.mark.unit +def test_emit_execution_decision_can_be_disabled_or_sampled_out(caplog): + decision = select_execution_decision(operator="linear_logp", stage="train_logprob") + + with caplog.at_level(logging.INFO, logger="vime.backends.rl_kernel_utils.execution"): + record = emit_execution_decision(decision, enabled=False) + sampled = emit_execution_decision(decision, sample_rate=0.5, random_value=0.75) + + assert record["decision"] == "native" + assert sampled["decision"] == "native" + assert caplog.records == [] diff --git a/vime/backends/rl_kernel_utils/__init__.py b/vime/backends/rl_kernel_utils/__init__.py new file mode 100644 index 00000000..dafb50eb --- /dev/null +++ b/vime/backends/rl_kernel_utils/__init__.py @@ -0,0 +1,27 @@ +from vime.backends.rl_kernel_utils.execution import ( + BackendCapability, + CapabilityQueryResult, + ExecutionDecision, + FallbackReason, + LogprobContractMetadata, + NumericContract, + RlKernelCapabilities, + build_logprob_contract_decision, + emit_execution_decision, + query_rl_kernel_capabilities, + select_execution_decision, +) + +__all__ = [ + "BackendCapability", + "CapabilityQueryResult", + "ExecutionDecision", + "FallbackReason", + "LogprobContractMetadata", + "NumericContract", + "RlKernelCapabilities", + "build_logprob_contract_decision", + "emit_execution_decision", + "query_rl_kernel_capabilities", + "select_execution_decision", +] diff --git a/vime/backends/rl_kernel_utils/execution.py b/vime/backends/rl_kernel_utils/execution.py new file mode 100644 index 00000000..fd2746ba --- /dev/null +++ b/vime/backends/rl_kernel_utils/execution.py @@ -0,0 +1,528 @@ +import json +import logging +from dataclasses import asdict, dataclass, field +from typing import Any + +logger = logging.getLogger(__name__) + +RLK_DECISION_EVENT = "rl_kernel.execution_decision" +_NATIVE_BACKEND_ID = "vime.native" + + +@dataclass(frozen=True) +class FallbackReason: + code: str + message: str + details: dict[str, Any] = field(default_factory=dict) + + +@dataclass(frozen=True) +class NumericContract: + contract_id: str + accumulation_dtype: str | None = None + reduction_order: str | None = None + downcast_point: str | None = None + sharding_rule: str | None = None + merge_semantics: str | None = None + quantization_policy: str | None = None + backend_id: str | None = None + tolerance_by_dtype: dict[str, Any] = field(default_factory=dict) + + def tolerance_for_dtype(self, dtype: str) -> Any: + try: + return self.tolerance_by_dtype[dtype] + except KeyError as exc: + raise KeyError(f"contract {self.contract_id!r} has no tolerance for dtype {dtype!r}") from exc + + +@dataclass(frozen=True) +class LogprobContractMetadata: + source: str + contract_id: str | None + dtype: str | None = None + backend_id: str | None = None + + +@dataclass(frozen=True) +class BackendCapability: + operator: str + backend_id: str + implementation_kind: str + dtypes: tuple[str, ...] = () + hardware_targets: tuple[str, ...] = () + autograd_modes: tuple[str, ...] = () + parallel_modes: tuple[str, ...] = () + shape_ranges: dict[str, Any] = field(default_factory=dict) + deterministic: bool = False + batch_invariant: bool = False + strict_fast_eligible: bool = False + production: bool = False + fallback_behavior: str | None = None + runtime_fingerprint: str | None = None + build_fingerprint: str | None = None + config_lifecycle: str | None = None + numeric_contract: NumericContract | None = None + + @property + def contract_id(self) -> str | None: + return None if self.numeric_contract is None else self.numeric_contract.contract_id + + @property + def is_reference(self) -> bool: + return self.implementation_kind == "reference" + + def supports_dtype(self, dtype: str | None) -> bool: + return dtype is None or not self.dtypes or dtype in self.dtypes + + +@dataclass(frozen=True) +class RlKernelCapabilities: + available: bool + backends: tuple[BackendCapability, ...] = () + reason: str | None = None + runtime_fingerprint: str | None = None + build_fingerprint: str | None = None + + def matching_backends(self, operator: str, dtype: str | None = None) -> tuple[BackendCapability, ...]: + return tuple( + backend for backend in self.backends if backend.operator == operator and backend.supports_dtype(dtype) + ) + + def tolerance_for_contract(self, contract_id: str, dtype: str) -> Any: + for backend in self.backends: + contract = backend.numeric_contract + if contract is not None and contract.contract_id == contract_id: + return contract.tolerance_for_dtype(dtype) + raise KeyError(f"unknown RL-Kernel numeric contract {contract_id!r}") + + +@dataclass(frozen=True) +class CapabilityQueryResult: + capabilities: RlKernelCapabilities + fallback_reason: FallbackReason | None = None + + +@dataclass(frozen=True) +class ExecutionDecision: + operator: str + stage: str + requested_mode: str + requested_backend: str | None + actual_backend: str | None + decision: str + fallback: bool = False + fallback_reason: FallbackReason | None = None + capability_backend_id: str | None = None + contract_id: str | None = None + dtype: str | None = None + parallel_context: dict[str, Any] = field(default_factory=dict) + strict_eligible: bool = False + details: dict[str, Any] = field(default_factory=dict) + + def to_log_record(self) -> dict[str, Any]: + return { + "event": RLK_DECISION_EVENT, + "operator": self.operator, + "stage": self.stage, + "requested_mode": self.requested_mode, + "requested_backend": self.requested_backend, + "actual_backend": self.actual_backend, + "decision": self.decision, + "fallback": self.fallback, + "fallback_reason": None if self.fallback_reason is None else asdict(self.fallback_reason), + "capability_backend_id": self.capability_backend_id, + "contract_id": self.contract_id, + "dtype": self.dtype, + "parallel_context": dict(self.parallel_context), + "strict_eligible": self.strict_eligible, + "details": dict(self.details), + } + + +def _native_contract_id(operator: str) -> str: + return f"{_NATIVE_BACKEND_ID}.{operator}" + + +def _normalize_caps(value: Any) -> RlKernelCapabilities: + if isinstance(value, RlKernelCapabilities): + return value + if isinstance(value, dict): + backends = tuple(_normalize_backend(backend) for backend in value.get("backends", ())) + return RlKernelCapabilities( + available=bool(value.get("available", True)), + backends=backends, + reason=value.get("reason"), + runtime_fingerprint=value.get("runtime_fingerprint"), + build_fingerprint=value.get("build_fingerprint"), + ) + raise TypeError(f"RL-Kernel capability provider returned unsupported value {type(value)!r}") + + +def _normalize_backend(value: Any) -> BackendCapability: + if isinstance(value, BackendCapability): + return value + if isinstance(value, dict): + data = dict(value) + contract = data.get("numeric_contract") + if isinstance(contract, dict): + data["numeric_contract"] = NumericContract(**contract) + for key in ("dtypes", "hardware_targets", "autograd_modes", "parallel_modes"): + if key in data and isinstance(data[key], list): + data[key] = tuple(data[key]) + return BackendCapability(**data) + raise TypeError(f"unsupported RL-Kernel backend descriptor {type(value)!r}") + + +def query_rl_kernel_capabilities(provider: Any = None) -> CapabilityQueryResult: + if provider is None: + reason = FallbackReason( + code="capability_provider_missing", + message="RL-Kernel capability provider was not configured.", + ) + return CapabilityQueryResult( + capabilities=RlKernelCapabilities(available=False, reason=reason.code), + fallback_reason=reason, + ) + + try: + if hasattr(provider, "get_vime_capabilities"): + raw_capabilities = provider.get_vime_capabilities() + elif callable(provider): + raw_capabilities = provider() + else: + raw_capabilities = provider + capabilities = _normalize_caps(raw_capabilities) + except Exception as exc: + reason = FallbackReason( + code="capability_query_failed", + message="Failed to query RL-Kernel capabilities.", + details={"error": repr(exc)}, + ) + return CapabilityQueryResult( + capabilities=RlKernelCapabilities(available=False, reason=reason.code), + fallback_reason=reason, + ) + + fallback_reason = None + if not capabilities.available: + fallback_reason = FallbackReason( + code=capabilities.reason or "rl_kernel_unavailable", + message="RL-Kernel reported that it is unavailable.", + ) + return CapabilityQueryResult(capabilities=capabilities, fallback_reason=fallback_reason) + + +def select_execution_decision( + *, + operator: str, + stage: str, + requested_fast: str = "off", + requested_consistency: str = "off", + capabilities: RlKernelCapabilities | None = None, + requested_backend: str | None = None, + dtype: str | None = None, + parallel_context: dict[str, Any] | None = None, +) -> ExecutionDecision: + requested_mode = f"fast={requested_fast},consistency={requested_consistency}" + parallel_context = parallel_context or {} + + if requested_fast == "off": + decision = "audit-only" if requested_consistency in {"audit", "strict"} else "native" + return ExecutionDecision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + actual_backend=_NATIVE_BACKEND_ID, + decision=decision, + contract_id=_native_contract_id(operator), + dtype=dtype, + parallel_context=parallel_context, + ) + + if capabilities is None or not capabilities.available: + reason = FallbackReason( + code="capability_data_missing", + message="RL-Kernel capability data is unavailable.", + ) + if requested_fast == "strict": + return _strict_failure_decision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + dtype=dtype, + parallel_context=parallel_context, + reason=reason, + ) + return _native_fallback_decision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + dtype=dtype, + parallel_context=parallel_context, + reason=reason, + ) + + candidates = _filter_requested_backend(capabilities.matching_backends(operator, dtype), requested_backend) + if requested_consistency == "strict": + strict_backend = _first_strict_fast_backend(candidates) + if strict_backend is not None: + return _backend_decision( + operator, + stage, + requested_mode, + requested_backend, + dtype, + parallel_context, + strict_backend, + "strict-fast", + ) + + reference_backend = _first_reference_backend(candidates) + if requested_fast == "auto" and reference_backend is not None: + return _backend_decision( + operator, + stage, + requested_mode, + requested_backend, + dtype, + parallel_context, + reference_backend, + "strict-reference", + ) + + reason = FallbackReason( + code="strict_backend_unavailable", + message="No RL-Kernel backend satisfies the requested strict consistency policy.", + details={"operator": operator, "requested_backend": requested_backend, "dtype": dtype}, + ) + return _strict_failure_decision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + dtype=dtype, + parallel_context=parallel_context, + reason=reason, + ) + + backend = _first_non_reference_backend(candidates) + if backend is not None: + return _backend_decision( + operator, stage, requested_mode, requested_backend, dtype, parallel_context, backend, "optimized" + ) + + reason = FallbackReason( + code="backend_unavailable", + message="No RL-Kernel backend matches the requested operator, backend, and dtype.", + details={"operator": operator, "requested_backend": requested_backend, "dtype": dtype}, + ) + if requested_fast == "strict": + return _strict_failure_decision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + dtype=dtype, + parallel_context=parallel_context, + reason=reason, + ) + return _native_fallback_decision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + dtype=dtype, + parallel_context=parallel_context, + reason=reason, + ) + + +def _filter_requested_backend( + candidates: tuple[BackendCapability, ...], + requested_backend: str | None, +) -> tuple[BackendCapability, ...]: + if requested_backend is None: + return candidates + return tuple(backend for backend in candidates if backend.backend_id == requested_backend) + + +def _first_strict_fast_backend(candidates: tuple[BackendCapability, ...]) -> BackendCapability | None: + return next((backend for backend in candidates if backend.strict_fast_eligible and not backend.is_reference), None) + + +def _first_reference_backend(candidates: tuple[BackendCapability, ...]) -> BackendCapability | None: + return next((backend for backend in candidates if backend.is_reference), None) + + +def _first_non_reference_backend(candidates: tuple[BackendCapability, ...]) -> BackendCapability | None: + return next((backend for backend in candidates if not backend.is_reference), None) + + +def _backend_decision( + operator: str, + stage: str, + requested_mode: str, + requested_backend: str | None, + dtype: str | None, + parallel_context: dict[str, Any], + backend: BackendCapability, + decision: str, +) -> ExecutionDecision: + return ExecutionDecision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + actual_backend=backend.backend_id, + decision=decision, + capability_backend_id=backend.backend_id, + contract_id=backend.contract_id, + dtype=dtype, + parallel_context=parallel_context, + strict_eligible=backend.strict_fast_eligible or backend.is_reference, + ) + + +def _native_fallback_decision( + *, + operator: str, + stage: str, + requested_mode: str, + requested_backend: str | None, + dtype: str | None, + parallel_context: dict[str, Any], + reason: FallbackReason, +) -> ExecutionDecision: + return ExecutionDecision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + actual_backend=_NATIVE_BACKEND_ID, + decision="fallback-native", + fallback=True, + fallback_reason=reason, + contract_id=_native_contract_id(operator), + dtype=dtype, + parallel_context=parallel_context, + ) + + +def _strict_failure_decision( + *, + operator: str, + stage: str, + requested_mode: str, + requested_backend: str | None, + dtype: str | None, + parallel_context: dict[str, Any], + reason: FallbackReason, +) -> ExecutionDecision: + return ExecutionDecision( + operator=operator, + stage=stage, + requested_mode=requested_mode, + requested_backend=requested_backend, + actual_backend=None, + decision="strict-failure", + fallback=False, + fallback_reason=reason, + dtype=dtype, + parallel_context=parallel_context, + ) + + +def build_logprob_contract_decision( + *, + rollout_contract: LogprobContractMetadata | None, + recomputed_contract: LogprobContractMetadata | None, + strict: bool, + operator: str = "logp", + stage: str = "audit", +) -> ExecutionDecision: + if rollout_contract is None or recomputed_contract is None: + reason = FallbackReason( + code="contract_id_missing", + message="Missing rollout or recomputed logprob contract ID.", + details={ + "rollout_contract_id": None if rollout_contract is None else rollout_contract.contract_id, + "recomputed_contract_id": None if recomputed_contract is None else recomputed_contract.contract_id, + }, + ) + return _contract_problem_decision(operator, stage, strict, reason) + + if rollout_contract.contract_id is None or recomputed_contract.contract_id is None: + reason = FallbackReason( + code="contract_id_missing", + message="Missing rollout or recomputed logprob contract ID.", + details={ + "rollout_contract_id": rollout_contract.contract_id, + "recomputed_contract_id": recomputed_contract.contract_id, + }, + ) + return _contract_problem_decision(operator, stage, strict, reason) + + if rollout_contract.contract_id != recomputed_contract.contract_id: + reason = FallbackReason( + code="contract_id_mismatch", + message="Rollout and recomputed logprob contract IDs do not match.", + details={ + "rollout_contract_id": rollout_contract.contract_id, + "recomputed_contract_id": recomputed_contract.contract_id, + }, + ) + return _contract_problem_decision(operator, stage, strict, reason) + + return ExecutionDecision( + operator=operator, + stage=stage, + requested_mode="contract-check", + requested_backend=rollout_contract.backend_id, + actual_backend=recomputed_contract.backend_id, + decision="contract-match", + contract_id=rollout_contract.contract_id, + dtype=recomputed_contract.dtype or rollout_contract.dtype, + strict_eligible=True, + details={ + "rollout_contract_id": rollout_contract.contract_id, + "recomputed_contract_id": recomputed_contract.contract_id, + }, + ) + + +def _contract_problem_decision( + operator: str, + stage: str, + strict: bool, + reason: FallbackReason, +) -> ExecutionDecision: + return ExecutionDecision( + operator=operator, + stage=stage, + requested_mode="contract-check", + requested_backend=None, + actual_backend=None if strict else _NATIVE_BACKEND_ID, + decision="strict-failure" if strict else "audit-warning", + fallback=False, + fallback_reason=reason, + details=reason.details, + ) + + +def emit_execution_decision( + decision: ExecutionDecision, + *, + log: logging.Logger | None = None, + enabled: bool = True, + sample_rate: float = 1.0, + random_value: float = 0.0, +) -> dict[str, Any]: + record = decision.to_log_record() + if not enabled or sample_rate <= 0 or random_value >= sample_rate: + return record + + (log or logger).info("%s %s", RLK_DECISION_EVENT, json.dumps(record, sort_keys=True)) + return record From 4cda449721f4be8fbf1d73bdb07b9962f2ecd5b2 Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Wed, 15 Jul 2026 00:07:09 +0800 Subject: [PATCH 2/5] Add RL-Kernel consistency metadata fingerprints Co-authored-by: OpenAI Codex Signed-off-by: inaniloquentee <3051000145@qq.com> --- tests/test_sample.py | 6 + tests/utils/test_consistency_metadata.py | 282 ++++++++ ...onsistency_metadata_rollout_integration.py | 106 +++ vime/ray/rollout.py | 27 + vime/rollout/vllm_rollout.py | 65 ++ vime/utils/consistency_metadata.py | 608 ++++++++++++++++++ vime/utils/types.py | 1 + 7 files changed, 1095 insertions(+) create mode 100644 tests/utils/test_consistency_metadata.py create mode 100644 tests/utils/test_consistency_metadata_rollout_integration.py create mode 100644 vime/utils/consistency_metadata.py diff --git a/tests/test_sample.py b/tests/test_sample.py index bc9d3ff6..d8bf082c 100644 --- a/tests/test_sample.py +++ b/tests/test_sample.py @@ -50,6 +50,11 @@ def _make_sample(**overrides) -> Sample: loss_mask=[1, 1, 0, 1, 1], weight_versions=["v1"], rollout_log_probs=[-0.1, -0.2], + consistency_metadata={ + "schema_version": 1, + "sample": {"rollout_id": 7, "index": 42}, + "fingerprint": "sha256:unit-test", + }, rollout_top_p_token_ids=[10, 11, 12, 20], rollout_top_p_token_offsets=[0, 3, 4], rollout_routed_experts=[[0, 1], [2, 3]], @@ -133,6 +138,7 @@ def test_round_trip_preserves_every_field(): "loss_mask", "weight_versions", "rollout_log_probs", + "consistency_metadata", "rollout_top_p_token_ids", "rollout_top_p_token_offsets", "rollout_routed_experts", diff --git a/tests/utils/test_consistency_metadata.py b/tests/utils/test_consistency_metadata.py new file mode 100644 index 00000000..fd1c6c81 --- /dev/null +++ b/tests/utils/test_consistency_metadata.py @@ -0,0 +1,282 @@ +from __future__ import annotations + +import argparse + +import pytest + +from vime.rollout.data_source import RolloutDataSourceWithBuffer +from vime.utils.consistency_metadata import ( + build_batch_layout_fingerprints, + build_requested_actual_provenance, + build_rollout_consistency_metadata, + ensure_sample_consistency_metadata, + get_consistency_mode, + raise_for_consistency_metadata_failures, + stable_fingerprint, + validate_samples_consistency_metadata, +) +from vime.utils.types import Sample + +NUM_GPUS = 0 + + +def _args(**overrides): + values = dict( + hf_checkpoint="unit/model", + padding_side="right", + num_gpus=2, + num_gpus_per_node=2, + rollout_num_gpus=2, + rollout_num_gpus_per_engine=1, + router_policy="consistent_hash", + vllm_router_ip="127.0.0.1", + vllm_router_port=8000, + rlk_consistency="audit", + buffer_filter_path=None, + rollout_global_dataset=False, + n_samples_per_prompt=2, + ) + values.update(overrides) + return argparse.Namespace(**values) + + +def _sample(**overrides) -> Sample: + values = dict( + index=5, + group_index=2, + rollout_id=11, + session_id="session-1", + tokens=[101, 201, 202, 203], + response_length=3, + loss_mask=[1, 0, 1], + rollout_log_probs=[-0.1, -0.2, -0.3], + weight_versions=["weights-v1"], + status=Sample.Status.COMPLETED, + metadata={}, + ) + values.update(overrides) + return Sample(**values) + + +def _complete_sample(**overrides) -> Sample: + sample = _sample(**overrides) + sample.consistency_metadata = build_rollout_consistency_metadata( + sample, + args=_args(), + sampling_params={"temperature": 0.7, "top_p": 0.95, "top_k": 50, "max_new_tokens": 128}, + logprob_contract_id="rlk.logp.fp32.v1", + requested_provenance={"backend": "native", "fallback": False}, + actual_provenance={"backend": "native", "fallback": False}, + ) + return sample + + +@pytest.mark.unit +def test_stable_fingerprint_is_order_insensitive_for_dicts(): + assert stable_fingerprint({"b": 2, "a": [1, 2]}) == stable_fingerprint({"a": [1, 2], "b": 2}) + assert stable_fingerprint({"a": [1, 2]}) != stable_fingerprint({"a": [2, 1]}) + + +@pytest.mark.unit +def test_rollout_metadata_records_compact_token_and_active_mask_fingerprints(): + sample = _complete_sample() + metadata = sample.consistency_metadata + + assert metadata["tokens"]["total_token_count"] == 4 + assert metadata["tokens"]["response_length"] == 3 + assert metadata["tokens"]["response_token_ids_fingerprint"] == stable_fingerprint([201, 202, 203]) + assert metadata["active_mask"]["active_token_count"] == 2 + assert metadata["active_mask"]["mask_fingerprint"] == stable_fingerprint([1, 0, 1]) + assert metadata["old_logp"]["source"] == "rollout_engine" + assert metadata["old_logp"]["contract_id"] == "rlk.logp.fp32.v1" + + +@pytest.mark.unit +def test_response_token_fingerprint_excludes_prompt_tokens(): + first = _complete_sample(tokens=[1, 10, 11], response_length=2, loss_mask=[1, 1]) + second = _complete_sample(tokens=[999, 10, 11], response_length=2, loss_mask=[1, 1]) + + assert ( + first.consistency_metadata["tokens"]["response_token_ids_fingerprint"] + == second.consistency_metadata["tokens"]["response_token_ids_fingerprint"] + ) + assert ( + first.consistency_metadata["tokens"]["token_ids_fingerprint"] + != second.consistency_metadata["tokens"]["token_ids_fingerprint"] + ) + + +@pytest.mark.unit +def test_metadata_uses_sample_index_as_default_rollout_identifier(): + sample = _complete_sample(rollout_id=None, index=42) + + assert sample.consistency_metadata["sample"]["rollout_id"] == 42 + validation = validate_samples_consistency_metadata([sample], mode="strict") + assert validation.ok + + +@pytest.mark.unit +def test_loss_mask_length_mismatch_is_rejected_before_audit_claims(): + sample = _sample(response_length=3, loss_mask=[1, 0]) + with pytest.raises(ValueError, match="loss_mask length"): + build_rollout_consistency_metadata( + sample, + args=_args(), + sampling_params={"temperature": 1.0}, + logprob_contract_id="contract", + ) + + +@pytest.mark.unit +def test_audit_mode_reports_missing_custom_metadata_as_structured_warning(): + validation = validate_samples_consistency_metadata([_sample(consistency_metadata=None)], mode="audit") + + assert validation.ok + assert [issue.code for issue in validation.warnings] == ["consistency_metadata_missing"] + assert validation.warnings[0].sample_index == 5 + assert validation.warnings[0].rollout_id == 11 + assert validation.failures == [] + + +@pytest.mark.unit +def test_strict_mode_fails_closed_when_required_metadata_is_missing(): + validation = validate_samples_consistency_metadata([_sample(consistency_metadata=None)], mode="strict") + + assert not validation.ok + assert [issue.code for issue in validation.failures] == ["consistency_metadata_missing"] + with pytest.raises(ValueError, match="Strict consistency metadata validation failed"): + raise_for_consistency_metadata_failures(validation) + + +@pytest.mark.unit +def test_complete_metadata_passes_strict_validation_and_counts_active_tokens(): + validation = validate_samples_consistency_metadata([_complete_sample()], mode="strict") + + assert validation.ok + assert validation.active_token_count == 2 + assert validation.zero_active_token_samples == [] + assert validation.warnings == [] + assert validation.failures == [] + + +@pytest.mark.unit +def test_zero_active_token_sample_is_identified_without_becoming_strict_failure(): + validation = validate_samples_consistency_metadata( + [_complete_sample(tokens=[1, 2], response_length=2, loss_mask=[0, 0])], + mode="strict", + ) + + assert validation.ok + assert validation.active_token_count == 0 + assert validation.zero_active_token_samples == [{"sample_index": 5, "rollout_id": 11}] + assert [issue.code for issue in validation.warnings] == ["zero_active_tokens"] + + +@pytest.mark.unit +def test_requested_actual_provenance_mismatch_is_audit_warning_and_strict_failure(): + sample = _complete_sample() + sample.consistency_metadata["provenance"] = build_requested_actual_provenance( + requested={"backend": "native", "fallback": False}, + actual={"backend": "fallback-native", "fallback": True}, + ) + + audit = validate_samples_consistency_metadata([sample], mode="audit") + assert {issue.code for issue in audit.warnings} == { + "requested_actual_provenance_mismatch", + "undeclared_runtime_fallback", + } + + strict = validate_samples_consistency_metadata([sample], mode="strict") + assert {issue.code for issue in strict.failures} == { + "requested_actual_provenance_mismatch", + "undeclared_runtime_fallback", + } + + +@pytest.mark.unit +def test_ensure_sample_consistency_metadata_does_not_overwrite_existing_record_by_default(): + sample = _sample(consistency_metadata={"schema_version": 1, "fingerprint": "existing"}) + + record = ensure_sample_consistency_metadata( + sample, + args=_args(), + sampling_params={"temperature": 1.0}, + logprob_contract_id="contract", + ) + + assert record == {"schema_version": 1, "fingerprint": "existing"} + assert sample.consistency_metadata == {"schema_version": 1, "fingerprint": "existing"} + + +@pytest.mark.unit +def test_batch_layout_fingerprints_map_samples_to_rank_and_microbatch(): + train_data = { + "tokens": [[1, 2], [3, 4, 5], [6, 7], [8, 9]], + "total_lengths": [2, 3, 2, 2], + "response_lengths": [2, 2, 2, 2], + "loss_masks": [[1, 1], [1, 0], [0, 1], [0, 0]], + "rollout_ids": [10, 11, 12, 13], + "sample_indices": [0, 1, 2, 3], + } + + layouts = build_batch_layout_fingerprints( + train_data, + partitions=[[0, 2], [1, 3]], + micro_batch_indices=[[[0], [1]], [[0, 1]]], + num_microbatches=[2], + global_batch_sizes=[2], + ) + + assert layouts[2]["dp_rank"] == 0 + assert layouts[2]["microbatch_id"] == 1 + assert layouts[2]["active_token_count"] == 1 + assert layouts[3]["dp_rank"] == 1 + assert layouts[3]["microbatch_size"] == 2 + assert layouts[3]["active_mask_density"] == 0.0 + assert layouts[0]["global_batch_shape"]["packed_order_fingerprint"] == stable_fingerprint([0, 2, 1, 3]) + + +@pytest.mark.unit +def test_batch_layout_fingerprints_mark_samples_dropped_by_existing_schedule_trim(): + train_data = { + "tokens": [[1, 2], [3, 4]], + "total_lengths": [2, 2], + "response_lengths": [2, 2], + "loss_masks": [[1, 1], [1, 0]], + "rollout_ids": [10, 11], + "sample_indices": [0, 1], + } + + layouts = build_batch_layout_fingerprints( + train_data, + partitions=[[0]], + micro_batch_indices=[[[0]]], + num_microbatches=[1], + global_batch_sizes=[1], + ) + + assert "dropped_by_schedule" not in layouts[0] + assert layouts[1]["dropped_by_schedule"] is True + assert layouts[1]["active_token_count"] == 1 + + +@pytest.mark.unit +def test_rollout_data_source_buffer_preserves_consistency_metadata(): + source = RolloutDataSourceWithBuffer(_args()) + sample = _complete_sample() + + source.add_samples([[sample, _complete_sample(index=6, rollout_id=12)]]) + returned = source.get_samples(1) + + assert returned[0][0].consistency_metadata == sample.consistency_metadata + + +@pytest.mark.unit +def test_get_consistency_mode_accepts_args_and_env_alias(monkeypatch): + assert get_consistency_mode(_args(rlk_consistency="strict")) == "strict" + + monkeypatch.setenv("VIME_RLK_CONSISTENCY", "audit") + assert get_consistency_mode(argparse.Namespace()) == "audit" + + with pytest.raises(ValueError, match="Unsupported RL-Kernel consistency mode"): + get_consistency_mode(_args(rlk_consistency="maybe")) diff --git a/tests/utils/test_consistency_metadata_rollout_integration.py b/tests/utils/test_consistency_metadata_rollout_integration.py new file mode 100644 index 00000000..146611d8 --- /dev/null +++ b/tests/utils/test_consistency_metadata_rollout_integration.py @@ -0,0 +1,106 @@ +from __future__ import annotations + +import sys +from pathlib import Path + +import pytest + +_tests_root = Path(__file__).resolve().parents[1] +if str(_tests_root) not in sys.path: + sys.path.insert(0, str(_tests_root)) + +import _unit_stubs + +_unit_stubs.install_rollout_optional_stubs() +_unit_stubs.install_vllm_cli_stubs() + +from vime.ray.rollout import RolloutManager # noqa: E402 +from vime.utils.consistency_metadata import build_rollout_consistency_metadata # noqa: E402 +from vime.utils.types import Sample # noqa: E402 + + +NUM_GPUS = 0 + + +class Args: + reward_key = None + advantage_estimator = "grpo" + rewards_normalization = False + grpo_std_normalization = False + rollout_top_p = 1.0 + rlk_consistency = "audit" + hf_checkpoint = "unit/model" + padding_side = "right" + num_gpus = 1 + rollout_num_gpus = 1 + rollout_num_gpus_per_engine = 1 + router_policy = "round_robin" + + +def _manager(mode: str): + cls = RolloutManager.__ray_actor_class__ + manager = cls.__new__(cls) + manager.args = Args() + manager.args.rlk_consistency = mode + manager.custom_reward_post_process_func = None + manager.custom_convert_samples_to_train_data_func = None + return manager + + +def _sample(index: int = 0) -> Sample: + return Sample( + index=index, + rollout_id=index, + tokens=[101, 201, 202], + response_length=2, + loss_mask=[1, 1], + rollout_log_probs=[-0.1, -0.2], + weight_versions=["w1"], + reward=1.0, + status=Sample.Status.COMPLETED, + metadata={}, + ) + + +@pytest.mark.unit +def test_convert_samples_to_train_data_leaves_consistency_fields_out_when_off(): + train_data = _manager("off")._convert_samples_to_train_data([_sample()]) + + assert "consistency_metadata" not in train_data + assert "consistency_metadata_validation" not in train_data + + +@pytest.mark.unit +def test_convert_samples_to_train_data_reports_audit_warning_for_missing_metadata(): + train_data = _manager("audit")._convert_samples_to_train_data([_sample()]) + + assert train_data["consistency_metadata"] == [None] + assert train_data["consistency_metadata_validation"]["ok"] is True + assert [issue["code"] for issue in train_data["consistency_metadata_validation"]["warnings"]] == [ + "consistency_metadata_missing" + ] + + +@pytest.mark.unit +def test_convert_samples_to_train_data_fails_closed_in_strict_mode(): + with pytest.raises(ValueError, match="Strict consistency metadata validation failed"): + _manager("strict")._convert_samples_to_train_data([_sample()]) + + +@pytest.mark.unit +def test_convert_samples_to_train_data_carries_sample_metadata_when_present(): + sample = _sample() + sample.consistency_metadata = build_rollout_consistency_metadata( + sample, + args=Args(), + sampling_params={"temperature": 1.0, "top_p": 1.0}, + logprob_contract_id="contract-v1", + requested_provenance={"backend": "native", "fallback": False}, + actual_provenance={"backend": "native", "fallback": False}, + ) + + train_data = _manager("strict")._convert_samples_to_train_data([sample]) + + assert train_data["consistency_metadata"] == [sample.consistency_metadata] + assert train_data["consistency_metadata_validation"]["ok"] is True + assert train_data["consistency_metadata_validation"]["active_token_count"] == 2 diff --git a/vime/ray/rollout.py b/vime/ray/rollout.py index 0607258d..0d1bd1e3 100644 --- a/vime/ray/rollout.py +++ b/vime/ray/rollout.py @@ -22,6 +22,13 @@ GPU_MEMORY_TYPE_CUDA_GRAPH = "cuda_graph" from vime.rollout.base_types import call_rollout_fn from vime.utils import logging_utils +from vime.utils.consistency_metadata import ( + build_batch_layout_fingerprints, + get_consistency_mode, + raise_for_consistency_metadata_failures, + sample_consistency_metadata, + validate_samples_consistency_metadata, +) from vime.utils.data import get_source from vime.utils.dp_schedule import build_dp_schedule from vime.utils.health_monitor import RolloutHealthMonitor @@ -752,6 +759,13 @@ def _convert_samples_to_train_data(self, samples: list[Sample] | list[list[Sampl loss_masks.append(sample.loss_mask) train_data["loss_masks"] = loss_masks + consistency_mode = get_consistency_mode(self.args) + if consistency_mode != "off": + validation = validate_samples_consistency_metadata(samples, mode=consistency_mode) + train_data["consistency_metadata_validation"] = validation.to_dict() + raise_for_consistency_metadata_failures(validation) + train_data["consistency_metadata"] = [sample_consistency_metadata(sample) for sample in samples] + # Per-rollout aggregate, precomputed at the step level (where we can # see every sample of every rollout) and broadcast per-sample so the # per-mb loss reducer uses the correct whole-rollout denominator even @@ -845,6 +859,15 @@ def _split_train_data_by_dp(self, data): rollout_indices=data["rollout_ids"], ) + if get_consistency_mode(self.args) != "off" or "consistency_metadata" in data: + data["consistency_batch_layout_fingerprints"] = build_batch_layout_fingerprints( + data, + partitions=partitions, + micro_batch_indices=micro_batch_indices, + num_microbatches=num_microbatches, + global_batch_sizes=global_batch_sizes, + ) + # Package per-rank rollout_data rollout_data_refs = [] for r in range(dp_size): @@ -868,6 +891,8 @@ def _split_train_data_by_dp(self, data): "source_names", "prompt", "teacher_log_probs", + "consistency_metadata", + "consistency_batch_layout_fingerprints", ]: if key not in data: continue @@ -877,6 +902,8 @@ def _split_train_data_by_dp(self, data): if key not in data: continue rollout_data[key] = data[key] + if "consistency_metadata_validation" in data: + rollout_data["consistency_metadata_validation"] = data["consistency_metadata_validation"] rollout_data["global_batch_sizes"] = global_batch_sizes rollout_data["num_microbatches"] = num_microbatches rollout_data["micro_batch_indices"] = micro_batch_indices[r] diff --git a/vime/rollout/vllm_rollout.py b/vime/rollout/vllm_rollout.py index 7e657e71..ed5cc593 100644 --- a/vime/rollout/vllm_rollout.py +++ b/vime/rollout/vllm_rollout.py @@ -19,6 +19,11 @@ from vime.rollout.base_types import RolloutFnEvalOutput, RolloutFnTrainOutput from vime.rollout.filter_hub.base_types import MetricGatherer, call_dynamic_filter from vime.utils.async_utils import run +from vime.utils.consistency_metadata import ( + ensure_sample_consistency_metadata, + get_consistency_mode, + stable_fingerprint, +) from vime.utils.data import Dataset from vime.utils.eval_config import EvalDatasetConfig from vime.utils.http_utils import get, get_rollout_num_engines, post @@ -44,6 +49,52 @@ _ABORT_RESWEEP_INTERVAL_S = 3.0 +def _iter_samples(node: Sample | list[Any]): + if isinstance(node, Sample): + yield node + return + if isinstance(node, list): + for item in node: + yield from _iter_samples(item) + + +def _refresh_consistency_fingerprint(record: dict[str, Any]) -> None: + payload = dict(record) + payload.pop("fingerprint", None) + record["fingerprint"] = stable_fingerprint(payload) + + +def _annotate_consistency_metadata( + args: Namespace, + node: Sample | list[Any], + *, + sampling_params: dict[str, Any], + source: str, +) -> None: + if get_consistency_mode(args) == "off": + return + for sample in _iter_samples(node): + old_logp_source = source if sample.rollout_log_probs is not None else None + ensure_sample_consistency_metadata( + sample, + args=args, + sampling_params=sampling_params, + model_name=getattr(args, "hf_checkpoint", None), + old_logp_source=old_logp_source, + ) + + +def _record_dynamic_sampling_decision(node: Sample | list[Any], *, keep: bool, reason: str | None) -> None: + decision = {"keep": bool(keep), "reason": reason} + for sample in _iter_samples(node): + if sample.metadata is None: + sample.metadata = {} + sample.metadata["dynamic_sampling"] = decision + if sample.consistency_metadata is not None: + sample.consistency_metadata["dynamic_sampling"] = decision + _refresh_consistency_fingerprint(sample.consistency_metadata) + + def _coerce_flat_int_token_ids(ids: Any) -> list[int]: """Flatten tokenizer/processor output into ``list[int]`` for vLLM ``/inference/v1/generate``.""" if ids is None: @@ -438,6 +489,7 @@ async def generate_and_rm( with state.dp_rank_context() as _: # Check sample.generate_function_path for per-sample custom_generate_function_path (e.g., from eval dataset config) custom_func_path = getattr(sample, "generate_function_path", None) or args.custom_generate_function_path + consistency_source = "custom_generate" if custom_func_path is not None else "vllm_rollout" if custom_func_path is not None: custom_generate_func = load_function(custom_func_path) @@ -449,6 +501,13 @@ async def generate_and_rm( else: sample = await generate(args, sample, sampling_params) + _annotate_consistency_metadata( + args, + sample, + sampling_params=sampling_params, + source=consistency_source, + ) + # for the rm that need the whole group, we will not do the rm here if args.group_rm: return sample @@ -627,6 +686,12 @@ async def generate_rollout_async( all_data.append(group) dynamic_filter_output = call_dynamic_filter(dynamic_filter, args, group) + if get_consistency_mode(args) != "off": + _record_dynamic_sampling_decision( + group, + keep=dynamic_filter_output.keep, + reason=dynamic_filter_output.reason, + ) if not dynamic_filter_output.keep: metric_gatherer.on_dynamic_filter_drop(reason=dynamic_filter_output.reason) state.remaining_batch_size -= 1 diff --git a/vime/utils/consistency_metadata.py b/vime/utils/consistency_metadata.py new file mode 100644 index 00000000..8fcc3f9d --- /dev/null +++ b/vime/utils/consistency_metadata.py @@ -0,0 +1,608 @@ +"""Compact consistency metadata for rollout-to-training audit paths. + +The helpers in this module keep the rollout boundary observable without +shipping full token or mask payloads inside metadata records. Full tensors +continue to live in the existing training batch fields; metadata carries +stable fingerprints, active-token counts, and requested-vs-actual provenance +needed by audit/strict consistency checks. +""" + +from __future__ import annotations + +import dataclasses +import hashlib +import json +import os +from dataclasses import dataclass, field +from enum import Enum +from typing import Any + +CONSISTENCY_METADATA_SCHEMA_VERSION = 1 +CONSISTENCY_MODES = {"off", "audit", "strict"} + +REQUIRED_COMPARISON_FIELDS = ( + ("sample.rollout_id", "rollout_id_missing"), + ("tokens.response_token_ids_fingerprint", "token_ids_missing"), + ("active_mask.mask_fingerprint", "active_mask_missing"), + ("active_mask.active_token_count", "active_mask_missing"), + ("tokenizer.fingerprint", "tokenizer_missing"), + ("sampling.params_fingerprint", "sampling_config_missing"), + ("padding.side", "padding_semantics_missing"), + ("model.name", "model_name_missing"), + ("weight.version", "weight_version_missing"), + ("old_logp.source", "old_logp_source_missing"), + ("old_logp.contract_id", "logprob_contract_id_missing"), + ("provenance.actual_fingerprint", "actual_provenance_missing"), +) + +_SAMPLING_PARAM_KEYS = ( + "temperature", + "top_p", + "top_k", + "max_new_tokens", + "max_tokens", + "seed", + "stop", + "stop_token_ids", + "skip_special_tokens", +) + +_PARALLEL_ARG_KEYS = ( + "num_gpus", + "num_gpus_per_node", + "rollout_num_gpus", + "rollout_num_gpus_per_engine", + "megatron_tensor_parallel_size", + "megatron_context_parallel_size", + "megatron_expert_model_parallel_size", + "tensor_model_parallel_size", + "context_parallel_size", +) + +_ROUTER_ARG_KEYS = ( + "router_policy", + "vllm_router_ip", + "vllm_router_port", + "vllm_dp_size", + "vllm_enable_prefix_caching", + "vllm_enable_deterministic_inference", + "rollout_data_transport", +) + + +@dataclass(frozen=True) +class ConsistencyMetadataIssue: + code: str + message: str + severity: str + sample_index: int | None = None + rollout_id: int | None = None + field: str | None = None + + def to_dict(self) -> dict[str, Any]: + return { + "code": self.code, + "message": self.message, + "severity": self.severity, + "sample_index": self.sample_index, + "rollout_id": self.rollout_id, + "field": self.field, + } + + +@dataclass(frozen=True) +class ConsistencyMetadataValidation: + mode: str + active_token_count: int = 0 + zero_active_token_samples: list[dict[str, int | None]] = field(default_factory=list) + warnings: list[ConsistencyMetadataIssue] = field(default_factory=list) + failures: list[ConsistencyMetadataIssue] = field(default_factory=list) + + @property + def ok(self) -> bool: + return not self.failures + + def to_dict(self) -> dict[str, Any]: + return { + "mode": self.mode, + "ok": self.ok, + "active_token_count": self.active_token_count, + "zero_active_token_samples": self.zero_active_token_samples, + "warnings": [issue.to_dict() for issue in self.warnings], + "failures": [issue.to_dict() for issue in self.failures], + } + + +def _normalize_for_json(value: Any) -> Any: + if dataclasses.is_dataclass(value): + return _normalize_for_json(dataclasses.asdict(value)) + if isinstance(value, Enum): + return value.value + if isinstance(value, dict): + return {str(key): _normalize_for_json(value[key]) for key in sorted(value, key=str)} + if isinstance(value, (list, tuple)): + return [_normalize_for_json(item) for item in value] + if hasattr(value, "detach") and callable(value.detach): + value = value.detach().cpu() + if hasattr(value, "tolist") and callable(value.tolist) and not isinstance(value, (str, bytes)): + return _normalize_for_json(value.tolist()) + if isinstance(value, bytes): + return value.hex() + if value is None or isinstance(value, (bool, int, float, str)): + return value + return repr(value) + + +def stable_fingerprint(value: Any) -> str: + payload = json.dumps( + _normalize_for_json(value), + sort_keys=True, + separators=(",", ":"), + ensure_ascii=True, + ) + return "sha256:" + hashlib.sha256(payload.encode("utf-8")).hexdigest() + + +def get_consistency_mode(args: Any | None = None) -> str: + value = None + if args is not None: + value = ( + getattr(args, "rlk_consistency", None) + or getattr(args, "rl_kernel_consistency", None) + or getattr(args, "rl_kernel_consistency_mode", None) + ) + if value is None: + value = os.environ.get("VIME_RLK_CONSISTENCY") or os.environ.get("VIME_RL_KERNEL_CONSISTENCY") + mode = str(value or "off").lower() + if mode not in CONSISTENCY_MODES: + raise ValueError( + f"Unsupported RL-Kernel consistency mode {value!r}; expected one of {sorted(CONSISTENCY_MODES)}" + ) + return mode + + +def _to_list(value: Any | None) -> list[Any]: + if value is None: + return [] + if hasattr(value, "detach") and callable(value.detach): + value = value.detach().cpu() + if hasattr(value, "tolist") and callable(value.tolist) and not isinstance(value, (str, bytes)): + value = value.tolist() + return list(value) + + +def _to_int_mask(mask: Any | None, response_length: int) -> list[int]: + if mask is None: + return [1] * response_length + values = [int(v) for v in _to_list(mask)] + if len(values) != response_length: + raise ValueError(f"loss_mask length {len(values)} != response_length {response_length}") + return values + + +def count_active_tokens(loss_mask: Any | None, response_length: int) -> int: + return sum(1 for value in _to_int_mask(loss_mask, response_length) if value) + + +def _metadata_dict(sample: Any) -> dict[str, Any]: + metadata = getattr(sample, "metadata", None) + return metadata if isinstance(metadata, dict) else {} + + +def _sample_status_value(sample: Any) -> str | None: + status = getattr(sample, "status", None) + return getattr(status, "value", status) + + +def _compact_attrs(obj: Any | None, keys: tuple[str, ...]) -> dict[str, Any]: + if obj is None: + return {} + return {key: getattr(obj, key) for key in keys if hasattr(obj, key) and getattr(obj, key) is not None} + + +def _compact_sampling_params(sampling_params: dict[str, Any] | None) -> dict[str, Any]: + if not sampling_params: + return {} + return {key: sampling_params[key] for key in _SAMPLING_PARAM_KEYS if key in sampling_params} + + +def build_requested_actual_provenance( + *, + requested: dict[str, Any] | None = None, + actual: dict[str, Any] | None = None, +) -> dict[str, Any]: + requested = _normalize_for_json(requested or {}) + actual = _normalize_for_json(actual or {}) + mismatches = {} + for key in sorted(set(requested) | set(actual)): + if key in requested and key in actual and requested[key] != actual[key]: + mismatches[key] = {"requested": requested[key], "actual": actual[key]} + + fallback = actual.get("fallback") + requested_fallback = requested.get("fallback") + undeclared_fallback = bool(fallback) and requested_fallback is not True + + return { + "requested": requested, + "actual": actual, + "requested_fingerprint": stable_fingerprint(requested) if requested else None, + "actual_fingerprint": stable_fingerprint(actual) if actual else None, + "mismatches": mismatches, + "undeclared_fallback": undeclared_fallback, + } + + +def _extract_requested_provenance(args: Any | None) -> dict[str, Any]: + requested = _compact_attrs(args, _ROUTER_ARG_KEYS + _PARALLEL_ARG_KEYS) + requested.update(_compact_attrs(args, ("hf_checkpoint", "load", "save"))) + return requested + + +def _extract_actual_provenance(args: Any | None, *, model_name: str | None) -> dict[str, Any]: + actual = _compact_attrs(args, _ROUTER_ARG_KEYS + _PARALLEL_ARG_KEYS) + if model_name is not None: + actual["model_name"] = model_name + return actual + + +def build_rollout_consistency_metadata( + sample: Any, + *, + args: Any | None = None, + sampling_params: dict[str, Any] | None = None, + model_name: str | None = None, + old_logp_source: str | None = None, + logprob_contract_id: str | None = None, + tokenizer_fingerprint: str | None = None, + padding_side: str | None = None, + requested_provenance: dict[str, Any] | None = None, + actual_provenance: dict[str, Any] | None = None, + batch_layout: dict[str, Any] | None = None, + dynamic_sampling: dict[str, Any] | None = None, +) -> dict[str, Any]: + metadata = _metadata_dict(sample) + response_length = int(getattr(sample, "response_length", 0) or 0) + tokens = [int(token) for token in _to_list(getattr(sample, "tokens", []))] + response_tokens = tokens[-response_length:] if response_length else [] + active_mask = _to_int_mask(getattr(sample, "loss_mask", None), response_length) + active_token_count = sum(1 for value in active_mask if value) + rollout_log_probs = getattr(sample, "rollout_log_probs", None) + weight_versions = list(getattr(sample, "weight_versions", None) or []) + rollout_id = getattr(sample, "rollout_id", None) + if rollout_id is None: + rollout_id = getattr(sample, "index", None) + + sampling_summary = _compact_sampling_params(sampling_params) or metadata.get("sampling_params") or {} + model_name = model_name or metadata.get("model_name") or getattr(args, "hf_checkpoint", None) + old_logp_source = old_logp_source or metadata.get("old_logp_source") + if old_logp_source is None and rollout_log_probs is not None: + old_logp_source = "rollout_engine" + logprob_contract_id = logprob_contract_id or metadata.get("logprob_contract_id") + tokenizer_fingerprint = tokenizer_fingerprint or metadata.get("tokenizer_fingerprint") + if tokenizer_fingerprint is None and getattr(args, "hf_checkpoint", None) is not None: + tokenizer_fingerprint = stable_fingerprint({"hf_checkpoint": args.hf_checkpoint}) + padding_side = padding_side or metadata.get("padding_side") or getattr(args, "padding_side", None) + + position_cache = metadata.get("position_cache") or metadata.get("position_cache_metadata") + quantization = metadata.get("quantization") or _compact_attrs(args, ("quantization", "vllm_quantization")) + parallel_placement = metadata.get("parallel_placement") or _compact_attrs(args, _PARALLEL_ARG_KEYS) + dynamic_sampling = dynamic_sampling or metadata.get("dynamic_sampling") + + requested = requested_provenance if requested_provenance is not None else _extract_requested_provenance(args) + actual = ( + actual_provenance if actual_provenance is not None else _extract_actual_provenance(args, model_name=model_name) + ) + + record = { + "schema_version": CONSISTENCY_METADATA_SCHEMA_VERSION, + "sample": { + "group_index": getattr(sample, "group_index", None), + "index": getattr(sample, "index", None), + "rollout_id": rollout_id, + "session_id": getattr(sample, "session_id", None), + "status": _sample_status_value(sample), + }, + "tokens": { + "total_token_count": len(tokens), + "response_length": response_length, + "token_ids_fingerprint": stable_fingerprint(tokens), + "response_token_ids_fingerprint": stable_fingerprint(response_tokens), + }, + "active_mask": { + "response_length": response_length, + "active_token_count": active_token_count, + "mask_fingerprint": stable_fingerprint(active_mask), + "zero_active_tokens": active_token_count == 0, + }, + "tokenizer": {"fingerprint": tokenizer_fingerprint}, + "sampling": { + "params_fingerprint": stable_fingerprint(sampling_summary) if sampling_summary else None, + "summary": _normalize_for_json(sampling_summary), + }, + "padding": {"side": padding_side}, + "position_cache": { + "fingerprint": stable_fingerprint(position_cache) if position_cache else None, + }, + "quantization": { + "fingerprint": stable_fingerprint(quantization) if quantization else None, + "summary": _normalize_for_json(quantization), + }, + "parallel_placement": { + "fingerprint": stable_fingerprint(parallel_placement) if parallel_placement else None, + "summary": _normalize_for_json(parallel_placement), + }, + "model": {"name": model_name}, + "old_logp": { + "source": old_logp_source, + "contract_id": logprob_contract_id, + "num_values": len(rollout_log_probs) if rollout_log_probs is not None else None, + "fingerprint": stable_fingerprint(rollout_log_probs) if rollout_log_probs is not None else None, + }, + "weight": { + "version": metadata.get("weight_version") or (weight_versions[-1] if weight_versions else None), + "versions_fingerprint": stable_fingerprint(weight_versions) if weight_versions else None, + "pre_update": metadata.get("pre_update"), + }, + "provenance": build_requested_actual_provenance(requested=requested, actual=actual), + "batch_layout": batch_layout, + "dynamic_sampling": _normalize_for_json(dynamic_sampling), + } + record["fingerprint"] = stable_fingerprint(record) + return record + + +def ensure_sample_consistency_metadata( + sample: Any, + *, + args: Any | None = None, + sampling_params: dict[str, Any] | None = None, + overwrite: bool = False, + **metadata_kwargs: Any, +) -> dict[str, Any]: + existing = getattr(sample, "consistency_metadata", None) + if existing is not None and not overwrite: + return existing + record = build_rollout_consistency_metadata( + sample, + args=args, + sampling_params=sampling_params, + **metadata_kwargs, + ) + sample.consistency_metadata = record + return record + + +def sample_consistency_metadata(sample: Any) -> dict[str, Any] | None: + direct = getattr(sample, "consistency_metadata", None) + if direct is not None: + return direct + metadata = _metadata_dict(sample) + nested = metadata.get("consistency_metadata") + return nested if isinstance(nested, dict) else None + + +def _get_path(mapping: dict[str, Any], path: str) -> Any: + current: Any = mapping + for part in path.split("."): + if not isinstance(current, dict) or part not in current: + return None + current = current[part] + return current + + +def _sample_ids(sample: Any, metadata: dict[str, Any] | None) -> tuple[int | None, int | None]: + sample_index = getattr(sample, "index", None) + rollout_id = getattr(sample, "rollout_id", None) + if metadata: + sample_info = metadata.get("sample") or {} + sample_index = sample_info.get("index", sample_index) + rollout_id = sample_info.get("rollout_id", rollout_id) + if rollout_id is None: + rollout_id = sample_index + return sample_index, rollout_id + + +def _issue( + *, + mode: str, + code: str, + message: str, + sample_index: int | None, + rollout_id: int | None, + field: str | None = None, + strict_failure: bool = True, +) -> tuple[ConsistencyMetadataIssue, bool]: + is_failure = mode == "strict" and strict_failure + severity = "error" if is_failure else "warning" + return ( + ConsistencyMetadataIssue( + code=code, + message=message, + severity=severity, + sample_index=sample_index, + rollout_id=rollout_id, + field=field, + ), + is_failure, + ) + + +def validate_samples_consistency_metadata( + samples: list[Any], + *, + mode: str, + required_fields: tuple[tuple[str, str], ...] = REQUIRED_COMPARISON_FIELDS, +) -> ConsistencyMetadataValidation: + mode = str(mode).lower() + if mode not in CONSISTENCY_MODES: + raise ValueError(f"Unsupported consistency mode {mode!r}") + if mode == "off": + return ConsistencyMetadataValidation(mode=mode) + + warnings: list[ConsistencyMetadataIssue] = [] + failures: list[ConsistencyMetadataIssue] = [] + zero_active_token_samples: list[dict[str, int | None]] = [] + active_token_count = 0 + + def add_issue(issue: ConsistencyMetadataIssue, is_failure: bool) -> None: + if is_failure: + failures.append(issue) + else: + warnings.append(issue) + + for sample in samples: + metadata = sample_consistency_metadata(sample) + sample_index, rollout_id = _sample_ids(sample, metadata) + if metadata is None: + issue, is_failure = _issue( + mode=mode, + code="consistency_metadata_missing", + message="Sample is missing consistency metadata required for audit/strict comparison.", + sample_index=sample_index, + rollout_id=rollout_id, + ) + add_issue(issue, is_failure) + continue + + for path, code in required_fields: + if _get_path(metadata, path) in (None, ""): + issue, is_failure = _issue( + mode=mode, + code=code, + message=f"Consistency metadata field {path!r} is missing.", + sample_index=sample_index, + rollout_id=rollout_id, + field=path, + ) + add_issue(issue, is_failure) + + active = _get_path(metadata, "active_mask.active_token_count") + if active is not None: + active_token_count += int(active) + if int(active) == 0: + zero_active_token_samples.append({"sample_index": sample_index, "rollout_id": rollout_id}) + issue, is_failure = _issue( + mode=mode, + code="zero_active_tokens", + message="Sample has zero active response/action tokens for consistency aggregates.", + sample_index=sample_index, + rollout_id=rollout_id, + field="active_mask.active_token_count", + strict_failure=False, + ) + add_issue(issue, is_failure) + + provenance = metadata.get("provenance") if isinstance(metadata, dict) else None + if isinstance(provenance, dict): + mismatches = provenance.get("mismatches") or {} + if mismatches: + issue, is_failure = _issue( + mode=mode, + code="requested_actual_provenance_mismatch", + message="Requested-vs-actual provenance differs for compared sample.", + sample_index=sample_index, + rollout_id=rollout_id, + field="provenance.mismatches", + ) + add_issue(issue, is_failure) + if provenance.get("undeclared_fallback"): + issue, is_failure = _issue( + mode=mode, + code="undeclared_runtime_fallback", + message="Actual provenance reports fallback that was not declared in requested provenance.", + sample_index=sample_index, + rollout_id=rollout_id, + field="provenance.undeclared_fallback", + ) + add_issue(issue, is_failure) + + return ConsistencyMetadataValidation( + mode=mode, + active_token_count=active_token_count, + zero_active_token_samples=zero_active_token_samples, + warnings=warnings, + failures=failures, + ) + + +def raise_for_consistency_metadata_failures(validation: ConsistencyMetadataValidation) -> None: + if not validation.failures: + return + preview = "; ".join( + f"{issue.code}(sample_index={issue.sample_index}, rollout_id={issue.rollout_id}, field={issue.field})" + for issue in validation.failures[:5] + ) + remaining = len(validation.failures) - 5 + if remaining > 0: + preview += f"; ... {remaining} more" + raise ValueError(f"Strict consistency metadata validation failed: {preview}") + + +def build_batch_layout_fingerprints( + train_data: dict[str, Any], + *, + partitions: list[list[int]], + micro_batch_indices: list[list[list[int]]], + num_microbatches: list[int], + global_batch_sizes: list[int], +) -> list[dict[str, Any]]: + sample_count = len(train_data["tokens"]) + total_lengths = train_data["total_lengths"] + response_lengths = train_data["response_lengths"] + loss_masks = train_data["loss_masks"] + rollout_ids = train_data.get("rollout_ids", [None] * sample_count) + sample_indices = train_data.get("sample_indices", [None] * sample_count) + + packed_order = [sample_index for partition in partitions for sample_index in partition] + shared = { + "dp_size": len(partitions), + "num_microbatches": num_microbatches, + "global_batch_sizes": global_batch_sizes, + "sequence_lengths_fingerprint": stable_fingerprint(total_lengths), + "packed_order_fingerprint": stable_fingerprint(packed_order), + } + + layouts: list[dict[str, Any] | None] = [None] * sample_count + for dp_rank, partition in enumerate(partitions): + for microbatch_id, local_indices in enumerate(micro_batch_indices[dp_rank]): + for microbatch_offset, local_index in enumerate(local_indices): + global_index = partition[local_index] + response_length = int(response_lengths[global_index]) + active_tokens = count_active_tokens(loss_masks[global_index], response_length) + layout = { + "schema_version": CONSISTENCY_METADATA_SCHEMA_VERSION, + "sample_index": sample_indices[global_index], + "rollout_id": rollout_ids[global_index], + "total_length": int(total_lengths[global_index]), + "response_length": response_length, + "active_token_count": active_tokens, + "active_mask_density": active_tokens / response_length if response_length else 0.0, + "dp_rank": dp_rank, + "rank_local_index": local_index, + "microbatch_id": microbatch_id, + "microbatch_offset": microbatch_offset, + "microbatch_size": len(local_indices), + "global_batch_shape": shared, + } + layout["fingerprint"] = stable_fingerprint(layout) + layouts[global_index] = layout + + for global_index, layout in enumerate(layouts): + if layout is not None: + continue + response_length = int(response_lengths[global_index]) + active_tokens = count_active_tokens(loss_masks[global_index], response_length) + dropped_layout = { + "schema_version": CONSISTENCY_METADATA_SCHEMA_VERSION, + "sample_index": sample_indices[global_index], + "rollout_id": rollout_ids[global_index], + "total_length": int(total_lengths[global_index]), + "response_length": response_length, + "active_token_count": active_tokens, + "active_mask_density": active_tokens / response_length if response_length else 0.0, + "dropped_by_schedule": True, + "global_batch_shape": shared, + } + dropped_layout["fingerprint"] = stable_fingerprint(dropped_layout) + layouts[global_index] = dropped_layout + return [layout for layout in layouts if layout is not None] diff --git a/vime/utils/types.py b/vime/utils/types.py index 1e46c99c..d51255e3 100644 --- a/vime/utils/types.py +++ b/vime/utils/types.py @@ -119,6 +119,7 @@ class Sample: loss_mask: list[int] | None = None weight_versions: list[str] = field(default_factory=list) rollout_log_probs: list[float] | None = None # Log probabilities from rollout engine + consistency_metadata: dict[str, Any] | None = None # Ragged top-p nucleus token ids replayed from rollout sampling. For response # token i, kept ids are rollout_top_p_token_ids[offsets[i]:offsets[i + 1]]. rollout_top_p_token_ids: list[int] | torch.Tensor | None = None From 0ea526a382ec6026b5c97b228655c09fe69397bb Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Wed, 15 Jul 2026 17:56:26 +0800 Subject: [PATCH 3/5] Add audit-only dlogp diagnostics Co-authored-by: OpenAI Codex Signed-off-by: inaniloquentee <3051000145@qq.com> --- tests/utils/test_dlogp_diagnostics.py | 337 ++++++++++++++++++++ vime/backends/megatron_utils/data.py | 1 + vime/backends/megatron_utils/loss.py | 27 ++ vime/ray/rollout.py | 1 + vime/utils/arguments.py | 11 + vime/utils/dlogp_diagnostics.py | 433 ++++++++++++++++++++++++++ 6 files changed, 810 insertions(+) create mode 100644 tests/utils/test_dlogp_diagnostics.py create mode 100644 vime/utils/dlogp_diagnostics.py diff --git a/tests/utils/test_dlogp_diagnostics.py b/tests/utils/test_dlogp_diagnostics.py new file mode 100644 index 00000000..507ea8ed --- /dev/null +++ b/tests/utils/test_dlogp_diagnostics.py @@ -0,0 +1,337 @@ +from __future__ import annotations + +import importlib +import math +import sys +from argparse import Namespace +from pathlib import Path + +import pytest +import torch + +_tests_root = Path(__file__).resolve().parents[1] +if str(_tests_root) not in sys.path: + sys.path.insert(0, str(_tests_root)) + +import _unit_stubs + +from vime.utils.dlogp_diagnostics import compute_dlogp_diagnostics, get_rlk_consistency_mode, is_dlogp_audit_enabled + +NUM_GPUS = 0 + + +FULL_METADATA = { + "model_name": "unit-model", + "backend_id": "megatron", + "contract_id": "contract-v1", + "batch_layout_fingerprint": "layout-abc", + "provenance_fingerprint": "prov-def", +} + +LOSS_MODULE_PATH = "vime.backends.megatron_utils.loss" +MEGATRON_STUB_MODULES = ( + "megatron", + "megatron.core", + "megatron.core.parallel_state", + "megatron.core.transformer", + "megatron.core.transformer.transformer_layer", + LOSS_MODULE_PATH, +) + + +def _report_with_full_metadata(**kwargs): + return compute_dlogp_diagnostics( + model_name=FULL_METADATA["model_name"], + backend_id=FULL_METADATA["backend_id"], + contract_id=FULL_METADATA["contract_id"], + batch_layout_fingerprint=FULL_METADATA["batch_layout_fingerprint"], + provenance_fingerprint=FULL_METADATA["provenance_fingerprint"], + **kwargs, + ) + + +@pytest.fixture() +def megatron_loss_module(): + saved = _unit_stubs.save_sys_modules(MEGATRON_STUB_MODULES) + for module_name in MEGATRON_STUB_MODULES: + sys.modules.pop(module_name, None) + _unit_stubs.install_megatron_mpu_stub() + try: + yield importlib.import_module(LOSS_MODULE_PATH) + finally: + _unit_stubs.restore_sys_modules(saved) + + +def _policy_args(mode: str) -> Namespace: + return Namespace( + use_rollout_logprobs=False, + rollout_top_p=1.0, + use_opsm=False, + advantage_estimator="grpo", + eps_clip=0.2, + eps_clip_high=0.2, + get_mismatch_metrics=False, + use_tis=False, + custom_pg_loss_reducer_function_path=None, + entropy_coef=0.0, + use_kl_loss=False, + rlk_consistency_mode=mode, + model_name=FULL_METADATA["model_name"], + train_backend=FULL_METADATA["backend_id"], + rlk_contract_id=FULL_METADATA["contract_id"], + rlk_batch_layout_fingerprint=FULL_METADATA["batch_layout_fingerprint"], + rlk_provenance_fingerprint=FULL_METADATA["provenance_fingerprint"], + ) + + +def _policy_batch() -> dict: + return { + "advantages": [torch.ones(2)], + "log_probs": [torch.zeros(2)], + "rollout_log_probs": [torch.zeros(2)], + "response_lengths": [2], + "total_lengths": [2], + "loss_masks": [torch.ones(2, dtype=torch.int32)], + "unconcat_tokens": [torch.tensor([1, 2])], + "sample_indices": [42], + "rollout_ids": [7], + } + + +@pytest.mark.unit +def test_dlogp_metrics_use_active_tokens_only_and_identify_worst_token(): + train_log_probs = [ + torch.tensor([-1.0, -2.0, -3.0]), + torch.tensor([-0.1, -0.2, 10.0]), + ] + rollout_log_probs = [ + torch.tensor([-1.0, -1.5, -4.0]), + torch.tensor([-0.6, -0.2, -10.0]), + ] + loss_masks = [torch.tensor([1, 0, 1]), torch.tensor([1, 1, 0])] + + report = _report_with_full_metadata( + train_log_probs=train_log_probs, + rollout_log_probs=rollout_log_probs, + loss_masks=loss_masks, + sample_indices=torch.tensor([10, 11]), + rollout_ids=torch.tensor([3, 4]), + rank=7, + eps_clip=0.2, + ) + + # Active dlogp values are [0.0, 1.0, 0.5, 0.0]. The masked 20.0 delta must not win. + metrics = report.metrics + assert metrics["rlk_audit_active_token_count"].item() == pytest.approx(4.0) + assert metrics["rlk_audit_mask_coverage"].item() == pytest.approx(4.0 / 6.0) + assert metrics["rlk_audit_dlogp_abs_mean"].item() == pytest.approx(0.375) + assert metrics["rlk_audit_dlogp_abs_max"].item() == pytest.approx(1.0) + assert metrics["rlk_audit_dlogp_abs_p50"].item() == pytest.approx(0.25) + assert metrics["rlk_audit_dlogp_abs_p90"].item() == pytest.approx(0.85) + assert metrics["rlk_audit_dlogp_abs_p99"].item() == pytest.approx(0.985) + assert metrics["rlk_audit_warning_count"].item() == pytest.approx(0.0) + + assert report.worst_token == { + "abs_dlogp": pytest.approx(1.0), + "dlogp": pytest.approx(1.0), + "sample_position": 0, + "sample_id": 10, + "sample_index": 10, + "rollout_id": 3, + "token_position": 2, + "rank": 7, + **FULL_METADATA, + } + assert metrics["rlk_audit_worst_sample_index"].item() == pytest.approx(10.0) + assert metrics["rlk_audit_worst_rollout_id"].item() == pytest.approx(3.0) + assert metrics["rlk_audit_worst_token_position"].item() == pytest.approx(2.0) + assert metrics["rlk_audit_worst_rank"].item() == pytest.approx(7.0) + + +@pytest.mark.unit +def test_dlogp_ratio_clipfrac_and_approx_kl_follow_issue_formulas(): + dlogp = torch.tensor([0.0, math.log(1.5), math.log(0.75), math.log(1.1)]) + report = _report_with_full_metadata( + train_log_probs=[dlogp], + rollout_log_probs=[torch.zeros_like(dlogp)], + loss_masks=[torch.ones_like(dlogp)], + eps_clip=0.2, + ) + + ratio0 = dlogp.exp() + expected_clipfrac = ((ratio0 - 1.0).abs() > 0.2).float().mean() + expected_approx_kl = (ratio0 - 1.0 - dlogp).mean() + + assert report.metrics["rlk_audit_ratio0_mean"].item() == pytest.approx(ratio0.mean().item()) + assert report.metrics["rlk_audit_clipfrac0"].item() == pytest.approx(expected_clipfrac.item()) + assert report.metrics["rlk_audit_approx_kl0"].item() == pytest.approx(expected_approx_kl.item()) + + +@pytest.mark.unit +def test_dlogp_zero_active_sample_is_reported_without_worst_token(): + report = _report_with_full_metadata( + train_log_probs=[torch.tensor([1.0, 2.0])], + rollout_log_probs=[torch.tensor([1.0, 2.0])], + loss_masks=[torch.tensor([0, 0])], + ) + + assert report.metrics["rlk_audit_active_token_count"].item() == pytest.approx(0.0) + assert report.metrics["rlk_audit_mask_coverage"].item() == pytest.approx(0.0) + assert report.metrics["rlk_audit_zero_active_sample_count"].item() == pytest.approx(1.0) + assert report.metrics["rlk_audit_worst_token_position"].item() == pytest.approx(-1.0) + assert report.worst_token is None + assert [warning.code for warning in report.warnings] == ["zero_active_tokens"] + + +@pytest.mark.unit +def test_dlogp_shape_mismatch_is_a_structured_warning_and_skips_sample(): + report = _report_with_full_metadata( + train_log_probs=[torch.tensor([1.0, 2.0])], + rollout_log_probs=[torch.tensor([1.0])], + loss_masks=[torch.tensor([1, 1])], + sample_indices=[123], + ) + + assert report.metrics["rlk_audit_active_token_count"].item() == pytest.approx(0.0) + assert report.metrics["rlk_audit_warning_count"].item() == pytest.approx(1.0) + assert report.warnings[0].code == "shape_mismatch" + assert report.warnings[0].sample_id == 123 + + +@pytest.mark.unit +def test_dlogp_missing_rollout_log_probs_is_a_structured_warning(): + report = _report_with_full_metadata( + train_log_probs=[torch.tensor([1.0, 2.0])], + rollout_log_probs=None, + loss_masks=[torch.tensor([1, 1])], + ) + + assert report.metrics["rlk_audit_active_token_count"].item() == pytest.approx(0.0) + assert report.warnings[0].code == "missing_rollout_log_probs" + assert report.warnings[0].field == "rollout_log_probs" + + +@pytest.mark.unit +def test_dlogp_missing_metadata_produces_structured_warnings(): + report = compute_dlogp_diagnostics( + train_log_probs=[torch.tensor([1.0])], + rollout_log_probs=[torch.tensor([0.0])], + loss_masks=[torch.tensor([1])], + ) + + warning_fields = {warning.field for warning in report.warnings if warning.code == "missing_metadata"} + assert warning_fields == { + "model_name", + "backend_id", + "contract_id", + "batch_layout_fingerprint", + "provenance_fingerprint", + } + assert report.metrics["rlk_audit_warning_count"].item() == pytest.approx(5.0) + + +@pytest.mark.unit +def test_dlogp_sample_metadata_can_supply_context_fields(): + metadata = [ + { + "model_name": "sample-model", + "backend_id": "sample-backend", + "contract_id": "sample-contract", + "batch_layout_fingerprint": "sample-layout", + "provenance_fingerprint": "sample-provenance", + } + ] + + report = compute_dlogp_diagnostics( + train_log_probs=[torch.tensor([2.0])], + rollout_log_probs=[torch.tensor([0.0])], + loss_masks=[torch.tensor([1])], + metadata=metadata, + ) + + assert not report.warnings + assert report.worst_token is not None + assert report.worst_token["model_name"] == "sample-model" + assert report.worst_token["backend_id"] == "sample-backend" + assert report.worst_token["contract_id"] == "sample-contract" + assert report.worst_token["batch_layout_fingerprint"] == "sample-layout" + assert report.worst_token["provenance_fingerprint"] == "sample-provenance" + + +@pytest.mark.unit +def test_dlogp_diagnostics_are_read_only_and_detached(): + train = torch.tensor([1.0, 2.0], requires_grad=True) + rollout = torch.tensor([0.5, 1.5]) + mask = torch.tensor([1, 0]) + train_before = train.detach().clone() + rollout_before = rollout.clone() + mask_before = mask.clone() + + report = _report_with_full_metadata( + train_log_probs=[train], + rollout_log_probs=[rollout], + loss_masks=[mask], + ) + + torch.testing.assert_close(train.detach(), train_before) + torch.testing.assert_close(rollout, rollout_before) + torch.testing.assert_close(mask, mask_before) + assert all(not metric.requires_grad for metric in report.metrics.values()) + + +@pytest.mark.unit +def test_rlk_consistency_mode_defaults_to_off_and_supports_args_env_and_aliases(): + assert get_rlk_consistency_mode(None, environ={}) == "off" + assert not is_dlogp_audit_enabled(Namespace(), environ={}) + + assert get_rlk_consistency_mode(Namespace(rlk_consistency_mode="audit"), environ={}) == "audit" + assert is_dlogp_audit_enabled(Namespace(rlk_consistency_mode="strict"), environ={}) + assert is_dlogp_audit_enabled(Namespace(rlk_consistency_mode="audit"), environ={"VIME_RLK_CONSISTENCY": "off"}) + assert get_rlk_consistency_mode(Namespace(), environ={"VIME_RLK_CONSISTENCY": "audit-only"}) == "audit" + assert get_rlk_consistency_mode(Namespace(), environ={"VIME_RL_KERNEL_CONSISTENCY": "true"}) == "audit" + + with pytest.raises(ValueError, match="Unsupported RL-Kernel consistency mode"): + get_rlk_consistency_mode(Namespace(rlk_consistency_mode="fast"), environ={}) + + +@pytest.mark.unit +def test_policy_loss_adds_audit_metrics_only_when_enabled_without_changing_loss(monkeypatch, megatron_loss_module): + train_log_probs = [torch.tensor([0.1, 0.3])] + + def fake_get_log_probs_and_entropy(*args, **kwargs): + return None, {"log_probs": train_log_probs, "entropy": [torch.zeros(2)]} + + def fake_compute_policy_loss(ppo_kl, advantages, eps_clip, eps_clip_high): + del advantages, eps_clip, eps_clip_high + return torch.ones_like(ppo_kl), torch.zeros_like(ppo_kl) + + monkeypatch.setattr(megatron_loss_module, "get_log_probs_and_entropy", fake_get_log_probs_and_entropy) + monkeypatch.setattr(megatron_loss_module, "compute_policy_loss", fake_compute_policy_loss) + + def reducer(tensor): + return tensor.mean() + + logits = torch.zeros(1, 2, 4) + off_loss, off_metrics = megatron_loss_module.policy_loss_function( + _policy_args("off"), + _policy_batch(), + logits, + reducer, + ) + audit_loss, audit_metrics = megatron_loss_module.policy_loss_function( + _policy_args("audit"), + _policy_batch(), + logits, + reducer, + ) + + torch.testing.assert_close(audit_loss, off_loss) + assert not any(key.startswith("rlk_audit_") for key in off_metrics) + assert audit_metrics["rlk_audit_active_token_count"].item() == pytest.approx(2.0) + assert audit_metrics["rlk_audit_dlogp_abs_mean"].item() == pytest.approx(0.2) + assert audit_metrics["rlk_audit_worst_sample_index"].item() == pytest.approx(42.0) + assert audit_metrics["rlk_audit_worst_rollout_id"].item() == pytest.approx(7.0) + + +if __name__ == "__main__": + raise SystemExit(pytest.main([__file__])) diff --git a/vime/backends/megatron_utils/data.py b/vime/backends/megatron_utils/data.py index b5ec48f6..c1f4cdf7 100644 --- a/vime/backends/megatron_utils/data.py +++ b/vime/backends/megatron_utils/data.py @@ -290,6 +290,7 @@ def log_rollout_data( "global_batch_sizes", "num_microbatches", "micro_batch_indices", + "metadata", "source_names", ]: continue diff --git a/vime/backends/megatron_utils/loss.py b/vime/backends/megatron_utils/loss.py index 2003dd59..c24b5323 100644 --- a/vime/backends/megatron_utils/loss.py +++ b/vime/backends/megatron_utils/loss.py @@ -10,6 +10,7 @@ from torch.utils.checkpoint import checkpoint from vime.utils.distributed_utils import distributed_masked_whiten +from vime.utils.dlogp_diagnostics import compute_dlogp_diagnostics, is_dlogp_audit_enabled from vime.utils.misc import load_function from vime.utils.ppo_utils import ( calculate_log_probs_and_entropy, @@ -38,6 +39,12 @@ ) +def _get_dist_rank_or_none() -> int | None: + if not dist.is_available() or not dist.is_initialized(): + return None + return dist.get_rank() + + def get_rollout_top_p_logprob_kwargs(args: Namespace, batch: dict[str, Any]) -> dict[str, Any]: if args.rollout_top_p == 1.0: return {} @@ -932,6 +939,7 @@ def policy_loss_function( ) log_probs = log_probs_and_entropy["log_probs"] + audit_train_log_probs = log_probs if not args.use_rollout_logprobs and not old_log_probs: old_log_probs = [log_prob.detach() for log_prob in log_probs] train_log_probs_for_tis = batch.get("log_probs") @@ -1083,6 +1091,24 @@ def policy_loss_function( log_probs_to_compare = log_probs if args.use_rollout_logprobs else old_log_probs train_rollout_logprob_abs_diff = sum_of_sample_mean((log_probs_to_compare - rollout_log_probs).abs()) + dlogp_audit_metrics = {} + if is_dlogp_audit_enabled(args): + dlogp_audit_metrics = compute_dlogp_diagnostics( + audit_train_log_probs, + batch.get("rollout_log_probs"), + batch["loss_masks"], + sample_indices=batch.get("sample_indices"), + rollout_ids=batch.get("rollout_ids"), + metadata=batch.get("metadata"), + rank=_get_dist_rank_or_none(), + model_name=getattr(args, "model_name", None), + backend_id=getattr(args, "train_backend", "megatron"), + contract_id=getattr(args, "rlk_contract_id", None), + batch_layout_fingerprint=getattr(args, "rlk_batch_layout_fingerprint", None), + provenance_fingerprint=getattr(args, "rlk_provenance_fingerprint", None), + eps_clip=getattr(args, "eps_clip", 0.2), + ).metrics + reported_loss = { "loss": loss.clone().detach(), "pg_loss": pg_loss.clone().detach(), @@ -1093,6 +1119,7 @@ def policy_loss_function( if train_rollout_logprob_abs_diff is not None: reported_loss["train_rollout_logprob_abs_diff"] = train_rollout_logprob_abs_diff.clone().detach() + reported_loss.update(dlogp_audit_metrics) if args.use_kl_loss: reported_loss["kl_loss"] = kl_loss.clone().detach() diff --git a/vime/ray/rollout.py b/vime/ray/rollout.py index ec88fbfd..7fb938c4 100644 --- a/vime/ray/rollout.py +++ b/vime/ray/rollout.py @@ -872,6 +872,7 @@ def _split_train_data_by_dp(self, data): "rollout_top_p_token_ids", "rollout_top_p_token_offsets", "rollout_routed_experts", + "metadata", "source_names", "prompt", "teacher_log_probs", diff --git a/vime/utils/arguments.py b/vime/utils/arguments.py index 25fd9cc9..0b462fc2 100644 --- a/vime/utils/arguments.py +++ b/vime/utils/arguments.py @@ -1047,6 +1047,17 @@ def add_algo_arguments(parser): "If not set, we will use the logprobs from the actor model." ), ) + parser.add_argument( + "--rlk-consistency-mode", + "--rl-kernel-consistency-mode", + dest="rlk_consistency_mode", + choices=["off", "audit", "strict"], + default=None, + help=( + "RL-Kernel consistency diagnostics mode. 'audit' records native read-only dlogp telemetry " + "without enabling RL-Kernel fast operators; 'strict' currently records the same diagnostics." + ), + ) # Off-Policy Correction using Importance Sampling: https://fengyao.notion.site/off-policy-rl parser.add_argument( "--use-tis", diff --git a/vime/utils/dlogp_diagnostics.py b/vime/utils/dlogp_diagnostics.py new file mode 100644 index 00000000..7ac3ec7a --- /dev/null +++ b/vime/utils/dlogp_diagnostics.py @@ -0,0 +1,433 @@ +from __future__ import annotations + +import os +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Any + +import torch + +RLK_CONSISTENCY_MODES = ("off", "audit", "strict") +RLK_CONSISTENCY_MODE_ATTRS = ( + "rlk_consistency_mode", + "rlk_consistency", + "rl_kernel_consistency_mode", + "rl_kernel_consistency", +) +RLK_CONSISTENCY_MODE_ENVS = ( + "VIME_RLK_CONSISTENCY_MODE", + "VIME_RLK_CONSISTENCY", + "VIME_RL_KERNEL_CONSISTENCY_MODE", + "VIME_RL_KERNEL_CONSISTENCY", +) + +MISSING_METADATA_FIELDS = ( + "model_name", + "backend_id", + "contract_id", + "batch_layout_fingerprint", + "provenance_fingerprint", +) + + +@dataclass(frozen=True) +class DlogpAuditWarning: + code: str + message: str + sample_position: int | None = None + sample_id: int | str | None = None + field: str | None = None + + +@dataclass(frozen=True) +class DlogpAuditReport: + metrics: dict[str, torch.Tensor] + warnings: tuple[DlogpAuditWarning, ...] + worst_token: dict[str, Any] | None + + +def get_rlk_consistency_mode(args: Any | None, environ: Mapping[str, str] | None = None) -> str: + """Return the requested RL-Kernel consistency mode. + + The helper deliberately accepts several attribute/env names so this + audit-only feature can coexist with older or stacked branches that used + slightly different flag names. + """ + + raw_mode: Any | None = None + if args is not None: + for attr in RLK_CONSISTENCY_MODE_ATTRS: + value = getattr(args, attr, None) + if value is not None: + raw_mode = value + break + + if raw_mode is None: + env = os.environ if environ is None else environ + for env_name in RLK_CONSISTENCY_MODE_ENVS: + value = env.get(env_name) + if value is not None: + raw_mode = value + break + + if raw_mode is None: + return "off" + + mode = str(raw_mode).strip().lower().replace("_", "-") + aliases = { + "0": "off", + "false": "off", + "no": "off", + "none": "off", + "1": "audit", + "true": "audit", + "yes": "audit", + "audit-only": "audit", + } + mode = aliases.get(mode, mode) + if mode not in RLK_CONSISTENCY_MODES: + raise ValueError( + f"Unsupported RL-Kernel consistency mode {raw_mode!r}; " + f"expected one of {', '.join(RLK_CONSISTENCY_MODES)}." + ) + return mode + + +def is_dlogp_audit_enabled(args: Any | None, environ: Mapping[str, str] | None = None) -> bool: + return get_rlk_consistency_mode(args, environ=environ) in {"audit", "strict"} + + +def compute_dlogp_diagnostics( + train_log_probs: Sequence[torch.Tensor] | torch.Tensor, + rollout_log_probs: Sequence[torch.Tensor] | torch.Tensor | None, + loss_masks: Sequence[torch.Tensor] | torch.Tensor, + *, + sample_indices: Sequence[int | torch.Tensor] | torch.Tensor | None = None, + rollout_ids: Sequence[int | torch.Tensor] | torch.Tensor | None = None, + metadata: Sequence[Mapping[str, Any] | None] | None = None, + rank: int | None = None, + model_name: str | None = None, + backend_id: str | None = None, + contract_id: str | None = None, + batch_layout_fingerprint: str | None = None, + provenance_fingerprint: str | None = None, + eps_clip: float = 0.2, + prefix: str = "rlk_audit_", +) -> DlogpAuditReport: + """Compute read-only rollout-training dlogp diagnostics over active tokens.""" + + tensors = _as_tensor_list(train_log_probs) + masks = _as_tensor_list(loss_masks) + rollout_tensors = _as_tensor_list(rollout_log_probs) if rollout_log_probs is not None else [] + device = _first_device(tensors, masks, rollout_tensors) + dtype = _first_floating_dtype(tensors, rollout_tensors) + + warnings: list[DlogpAuditWarning] = [] + if not rollout_tensors: + warnings.append( + DlogpAuditWarning( + code="missing_rollout_log_probs", + message="rollout_log_probs are required for dlogp diagnostics.", + field="rollout_log_probs", + ) + ) + + if len(tensors) != len(masks): + warnings.append( + DlogpAuditWarning( + code="sample_count_mismatch", + message=f"train_log_probs has {len(tensors)} samples but loss_masks has {len(masks)}.", + ) + ) + if rollout_tensors and len(tensors) != len(rollout_tensors): + warnings.append( + DlogpAuditWarning( + code="sample_count_mismatch", + message=f"train_log_probs has {len(tensors)} samples but rollout_log_probs has {len(rollout_tensors)}.", + field="rollout_log_probs", + ) + ) + + context_metadata = { + "model_name": model_name, + "backend_id": backend_id, + "contract_id": contract_id, + "batch_layout_fingerprint": batch_layout_fingerprint, + "provenance_fingerprint": provenance_fingerprint, + } + for field, value in context_metadata.items(): + if value is None and not _metadata_field_available(metadata, field): + warnings.append( + DlogpAuditWarning( + code="missing_metadata", + message=f"{field} is not available for dlogp diagnostics.", + field=field, + ) + ) + + dlogp_parts: list[torch.Tensor] = [] + abs_parts: list[torch.Tensor] = [] + ratio_parts: list[torch.Tensor] = [] + clip_parts: list[torch.Tensor] = [] + approx_kl_parts: list[torch.Tensor] = [] + total_token_count = 0 + active_token_count = 0 + zero_active_sample_count = 0 + worst_token: dict[str, Any] | None = None + worst_abs_value: torch.Tensor | None = None + + sample_count = min(len(tensors), len(masks), len(rollout_tensors)) + with torch.no_grad(): + for sample_position in range(sample_count): + train = tensors[sample_position].detach().flatten() + rollout = rollout_tensors[sample_position].detach().flatten() + mask = masks[sample_position].detach().flatten() + sample_id = _value_at(sample_indices, sample_position) + rollout_id = _value_at(rollout_ids, sample_position) + sample_meta = ( + metadata[sample_position] if metadata is not None and sample_position < len(metadata) else None + ) + + if train.numel() != rollout.numel() or train.numel() != mask.numel(): + warnings.append( + DlogpAuditWarning( + code="shape_mismatch", + message=( + f"sample {sample_position} has train_log_probs={train.numel()}, " + f"rollout_log_probs={rollout.numel()}, loss_masks={mask.numel()}." + ), + sample_position=sample_position, + sample_id=sample_id, + ) + ) + continue + + total_token_count += int(mask.numel()) + active_mask = mask.to(dtype=torch.bool) + sample_active_count = int(active_mask.sum().item()) + active_token_count += sample_active_count + if sample_active_count == 0: + zero_active_sample_count += 1 + warnings.append( + DlogpAuditWarning( + code="zero_active_tokens", + message=f"sample {sample_position} has no active response/action tokens.", + sample_position=sample_position, + sample_id=sample_id, + ) + ) + continue + + active_positions = active_mask.nonzero(as_tuple=False).flatten() + sample_dlogp = (train.to(dtype=dtype) - rollout.to(dtype=dtype))[active_mask] + finite_mask = torch.isfinite(sample_dlogp) + if not bool(finite_mask.all().item()): + dropped = int((~finite_mask).sum().item()) + warnings.append( + DlogpAuditWarning( + code="non_finite_dlogp", + message=f"sample {sample_position} has {dropped} non-finite active dlogp values.", + sample_position=sample_position, + sample_id=sample_id, + ) + ) + active_positions = active_positions[finite_mask] + sample_dlogp = sample_dlogp[finite_mask] + if sample_dlogp.numel() == 0: + continue + + sample_abs = sample_dlogp.abs() + sample_ratio = sample_dlogp.exp() + sample_clip = ((sample_ratio - 1.0).abs() > eps_clip).to(dtype=dtype) + sample_approx_kl = sample_ratio - 1.0 - sample_dlogp + + dlogp_parts.append(sample_dlogp) + abs_parts.append(sample_abs) + ratio_parts.append(sample_ratio) + clip_parts.append(sample_clip) + approx_kl_parts.append(sample_approx_kl) + + sample_worst_abs, sample_worst_active_index = sample_abs.max(dim=0) + if worst_abs_value is None or bool((sample_worst_abs > worst_abs_value).item()): + token_position = int(active_positions[int(sample_worst_active_index.item())].item()) + worst_abs_value = sample_worst_abs + worst_token = { + "abs_dlogp": float(sample_worst_abs.item()), + "dlogp": float(sample_dlogp[int(sample_worst_active_index.item())].item()), + "sample_position": sample_position, + "sample_id": sample_id, + "sample_index": sample_id, + "rollout_id": rollout_id, + "token_position": token_position, + "rank": rank, + "model_name": _metadata_value("model_name", sample_meta, model_name), + "backend_id": _metadata_value("backend_id", sample_meta, backend_id), + "contract_id": _metadata_value("contract_id", sample_meta, contract_id), + "batch_layout_fingerprint": _metadata_value( + "batch_layout_fingerprint", + sample_meta, + batch_layout_fingerprint, + ), + "provenance_fingerprint": _metadata_value( + "provenance_fingerprint", + sample_meta, + provenance_fingerprint, + ), + } + + if dlogp_parts: + abs_dlogp = torch.cat(abs_parts).to(device=device, dtype=dtype) + ratio0 = torch.cat(ratio_parts).to(device=device, dtype=dtype) + clipfrac0_values = torch.cat(clip_parts).to(device=device, dtype=dtype) + approx_kl0_values = torch.cat(approx_kl_parts).to(device=device, dtype=dtype) + quantiles = torch.quantile(abs_dlogp, torch.tensor([0.5, 0.9, 0.99], device=device, dtype=dtype)) + metrics = { + f"{prefix}dlogp_abs_mean": abs_dlogp.mean(), + f"{prefix}dlogp_abs_max": abs_dlogp.max(), + f"{prefix}dlogp_abs_p50": quantiles[0], + f"{prefix}dlogp_abs_p90": quantiles[1], + f"{prefix}dlogp_abs_p99": quantiles[2], + f"{prefix}ratio0_mean": ratio0.mean(), + f"{prefix}clipfrac0": clipfrac0_values.mean(), + f"{prefix}approx_kl0": approx_kl0_values.mean(), + } + else: + metrics = { + f"{prefix}dlogp_abs_mean": _zero(device, dtype), + f"{prefix}dlogp_abs_max": _zero(device, dtype), + f"{prefix}dlogp_abs_p50": _zero(device, dtype), + f"{prefix}dlogp_abs_p90": _zero(device, dtype), + f"{prefix}dlogp_abs_p99": _zero(device, dtype), + f"{prefix}ratio0_mean": _zero(device, dtype), + f"{prefix}clipfrac0": _zero(device, dtype), + f"{prefix}approx_kl0": _zero(device, dtype), + } + + metrics.update( + { + f"{prefix}active_token_count": torch.tensor(float(active_token_count), device=device, dtype=dtype), + f"{prefix}mask_coverage": torch.tensor( + float(active_token_count / total_token_count) if total_token_count else 0.0, + device=device, + dtype=dtype, + ), + f"{prefix}zero_active_sample_count": torch.tensor( + float(zero_active_sample_count), device=device, dtype=dtype + ), + f"{prefix}warning_count": torch.tensor(float(len(warnings)), device=device, dtype=dtype), + } + ) + + if worst_token is not None: + metrics.update( + { + f"{prefix}worst_abs_dlogp": torch.tensor(worst_token["abs_dlogp"], device=device, dtype=dtype), + f"{prefix}worst_dlogp": torch.tensor(worst_token["dlogp"], device=device, dtype=dtype), + f"{prefix}worst_sample_position": torch.tensor( + float(worst_token["sample_position"]), + device=device, + dtype=dtype, + ), + f"{prefix}worst_sample_index": torch.tensor( + float(worst_token["sample_index"]) if _is_number(worst_token["sample_index"]) else -1.0, + device=device, + dtype=dtype, + ), + f"{prefix}worst_rollout_id": torch.tensor( + float(worst_token["rollout_id"]) if _is_number(worst_token["rollout_id"]) else -1.0, + device=device, + dtype=dtype, + ), + f"{prefix}worst_token_position": torch.tensor( + float(worst_token["token_position"]), + device=device, + dtype=dtype, + ), + f"{prefix}worst_rank": torch.tensor( + float(rank) if rank is not None else -1.0, + device=device, + dtype=dtype, + ), + } + ) + else: + metrics.update( + { + f"{prefix}worst_abs_dlogp": _zero(device, dtype), + f"{prefix}worst_dlogp": _zero(device, dtype), + f"{prefix}worst_sample_position": torch.tensor(-1.0, device=device, dtype=dtype), + f"{prefix}worst_sample_index": torch.tensor(-1.0, device=device, dtype=dtype), + f"{prefix}worst_rollout_id": torch.tensor(-1.0, device=device, dtype=dtype), + f"{prefix}worst_token_position": torch.tensor(-1.0, device=device, dtype=dtype), + f"{prefix}worst_rank": torch.tensor( + float(rank) if rank is not None else -1.0, device=device, dtype=dtype + ), + } + ) + + metrics = {key: value.clone().detach() for key, value in metrics.items()} + return DlogpAuditReport(metrics=metrics, warnings=tuple(warnings), worst_token=worst_token) + + +def _as_tensor_list(value: Sequence[torch.Tensor] | torch.Tensor | None) -> list[torch.Tensor]: + if value is None: + return [] + if isinstance(value, torch.Tensor): + return [value] + return [torch.as_tensor(item) for item in value] + + +def _first_device(*groups: Sequence[torch.Tensor]) -> torch.device: + for group in groups: + for tensor in group: + return tensor.device + return torch.device("cpu") + + +def _first_floating_dtype(*groups: Sequence[torch.Tensor]) -> torch.dtype: + for group in groups: + for tensor in group: + if tensor.is_floating_point(): + return tensor.dtype + return torch.float32 + + +def _zero(device: torch.device, dtype: torch.dtype) -> torch.Tensor: + return torch.tensor(0.0, device=device, dtype=dtype) + + +def _value_at(values: Sequence[Any] | torch.Tensor | None, index: int) -> Any: + if values is None: + return None + if isinstance(values, torch.Tensor): + if index >= values.numel(): + return None + return _python_scalar(values.flatten()[index]) + if index >= len(values): + return None + value = values[index] + if isinstance(value, torch.Tensor): + return _python_scalar(value) + return value + + +def _python_scalar(value: torch.Tensor) -> Any: + if value.numel() != 1: + return value.detach().cpu().tolist() + return value.detach().cpu().item() + + +def _metadata_value(field: str, sample_meta: Mapping[str, Any] | None, fallback: Any) -> Any: + if sample_meta is not None and field in sample_meta: + return sample_meta[field] + return fallback + + +def _metadata_field_available(metadata: Sequence[Mapping[str, Any] | None] | None, field: str) -> bool: + if metadata is None: + return False + return any(sample_meta is not None and sample_meta.get(field) is not None for sample_meta in metadata) + + +def _is_number(value: Any) -> bool: + return isinstance(value, int | float) and not isinstance(value, bool) From 9c13b7cf5ec03fd8b0d79cafa0d6ab7839dd64c4 Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Wed, 22 Jul 2026 19:57:38 +0800 Subject: [PATCH 4/5] Address RL-Kernel execution review feedback --- tests/utils/test_rl_kernel_execution.py | 66 +++++++++++++++++++- vime/backends/rl_kernel_utils/__init__.py | 2 + vime/backends/rl_kernel_utils/execution.py | 71 +++++++++++++++++++--- 3 files changed, 129 insertions(+), 10 deletions(-) diff --git a/tests/utils/test_rl_kernel_execution.py b/tests/utils/test_rl_kernel_execution.py index b332aa70..5a899a3c 100644 --- a/tests/utils/test_rl_kernel_execution.py +++ b/tests/utils/test_rl_kernel_execution.py @@ -11,6 +11,7 @@ RlKernelCapabilities, build_logprob_contract_decision, emit_execution_decision, + execution_decision_sample_value, query_rl_kernel_capabilities, select_execution_decision, ) @@ -270,15 +271,46 @@ def test_query_capabilities_normalizes_dict_provider(): @pytest.mark.unit -def test_query_capabilities_provider_failure_is_structured(): +def test_query_capabilities_filters_unknown_backend_fields(): + result = query_rl_kernel_capabilities( + lambda: { + "available": True, + "backends": [ + { + "operator": "linear_logp", + "backend_id": "rlk.linear_logp.fast", + "implementation_kind": "optimized", + "dtypes": ["bf16"], + "unknown_backend_key": "ignored", + "numeric_contract": { + "contract_id": "rlk.linear_logp.fp32", + "tolerance_by_dtype": {"bf16": {"source": "provider"}}, + "unknown_contract_key": "ignored", + }, + } + ], + } + ) + + assert result.fallback_reason is None + assert result.capabilities.available is True + assert result.capabilities.backends[0].backend_id == "rlk.linear_logp.fast" + assert result.capabilities.backends[0].dtypes == ("bf16",) + assert result.capabilities.backends[0].numeric_contract.contract_id == "rlk.linear_logp.fp32" + + +@pytest.mark.unit +def test_query_capabilities_provider_failure_is_structured_and_debug_logged(caplog): def provider(): raise RuntimeError("boom") - result = query_rl_kernel_capabilities(provider) + with caplog.at_level(logging.DEBUG, logger="vime.backends.rl_kernel_utils.execution"): + result = query_rl_kernel_capabilities(provider) assert result.capabilities.available is False assert result.fallback_reason.code == "capability_query_failed" assert "boom" in result.fallback_reason.details["error"] + assert any(record.message == "Provider query failed" and record.exc_info is not None for record in caplog.records) @pytest.mark.unit @@ -417,3 +449,33 @@ def test_emit_execution_decision_can_be_disabled_or_sampled_out(caplog): assert record["decision"] == "native" assert sampled["decision"] == "native" assert caplog.records == [] + + +@pytest.mark.unit +def test_execution_decision_sample_value_is_rank_specific_and_stable(): + rank0 = execution_decision_sample_value(seed="run-1", rank=0, key="linear_logp") + rank1 = execution_decision_sample_value(seed="run-1", rank=1, key="linear_logp") + + assert 0 <= rank0 < 1 + assert 0 <= rank1 < 1 + assert rank0 == execution_decision_sample_value(seed="run-1", rank=0, key="linear_logp") + assert rank0 != rank1 + + +@pytest.mark.unit +def test_emit_execution_decision_can_sample_with_rank_specific_values(caplog): + decision = select_execution_decision(operator="linear_logp", stage="train_logprob") + sample_key = "same-event" + rank0 = execution_decision_sample_value(seed="run-1", rank=0, key=sample_key) + rank1 = execution_decision_sample_value(seed="run-1", rank=1, key=sample_key) + sample_rate = (rank0 + rank1) / 2 + + with caplog.at_level(logging.INFO, logger="vime.backends.rl_kernel_utils.execution"): + emit_execution_decision( + decision, sample_rate=sample_rate, sample_seed="run-1", sample_rank=0, sample_key=sample_key + ) + emit_execution_decision( + decision, sample_rate=sample_rate, sample_seed="run-1", sample_rank=1, sample_key=sample_key + ) + + assert len(caplog.records) == 1 diff --git a/vime/backends/rl_kernel_utils/__init__.py b/vime/backends/rl_kernel_utils/__init__.py index dafb50eb..329bae26 100644 --- a/vime/backends/rl_kernel_utils/__init__.py +++ b/vime/backends/rl_kernel_utils/__init__.py @@ -8,6 +8,7 @@ RlKernelCapabilities, build_logprob_contract_decision, emit_execution_decision, + execution_decision_sample_value, query_rl_kernel_capabilities, select_execution_decision, ) @@ -22,6 +23,7 @@ "RlKernelCapabilities", "build_logprob_contract_decision", "emit_execution_decision", + "execution_decision_sample_value", "query_rl_kernel_capabilities", "select_execution_decision", ] diff --git a/vime/backends/rl_kernel_utils/execution.py b/vime/backends/rl_kernel_utils/execution.py index fd2746ba..90d87f48 100644 --- a/vime/backends/rl_kernel_utils/execution.py +++ b/vime/backends/rl_kernel_utils/execution.py @@ -1,13 +1,27 @@ +import hashlib import json import logging -from dataclasses import asdict, dataclass, field -from typing import Any +import os +from dataclasses import asdict, dataclass, field, fields +from typing import Any, Literal logger = logging.getLogger(__name__) RLK_DECISION_EVENT = "rl_kernel.execution_decision" _NATIVE_BACKEND_ID = "vime.native" +ExecutionDecisionKind = Literal[ + "native", + "audit-only", + "optimized", + "strict-fast", + "strict-reference", + "fallback-native", + "strict-failure", + "contract-match", + "audit-warning", +] + @dataclass(frozen=True) class FallbackReason: @@ -109,7 +123,7 @@ class ExecutionDecision: requested_mode: str requested_backend: str | None actual_backend: str | None - decision: str + decision: ExecutionDecisionKind fallback: bool = False fallback_reason: FallbackReason | None = None capability_backend_id: str | None = None @@ -162,10 +176,10 @@ def _normalize_backend(value: Any) -> BackendCapability: if isinstance(value, BackendCapability): return value if isinstance(value, dict): - data = dict(value) + data = _filter_dataclass_fields(value, BackendCapability) contract = data.get("numeric_contract") if isinstance(contract, dict): - data["numeric_contract"] = NumericContract(**contract) + data["numeric_contract"] = NumericContract(**_filter_dataclass_fields(contract, NumericContract)) for key in ("dtypes", "hardware_targets", "autograd_modes", "parallel_modes"): if key in data and isinstance(data[key], list): data[key] = tuple(data[key]) @@ -173,6 +187,11 @@ def _normalize_backend(value: Any) -> BackendCapability: raise TypeError(f"unsupported RL-Kernel backend descriptor {type(value)!r}") +def _filter_dataclass_fields(data: dict[str, Any], target: type[Any]) -> dict[str, Any]: + allowed = {field.name for field in fields(target)} + return {key: value for key, value in data.items() if key in allowed} + + def query_rl_kernel_capabilities(provider: Any = None) -> CapabilityQueryResult: if provider is None: reason = FallbackReason( @@ -193,6 +212,7 @@ def query_rl_kernel_capabilities(provider: Any = None) -> CapabilityQueryResult: raw_capabilities = provider capabilities = _normalize_caps(raw_capabilities) except Exception as exc: + logger.debug("Provider query failed", exc_info=True) reason = FallbackReason( code="capability_query_failed", message="Failed to query RL-Kernel capabilities.", @@ -369,7 +389,7 @@ def _backend_decision( dtype: str | None, parallel_context: dict[str, Any], backend: BackendCapability, - decision: str, + decision: ExecutionDecisionKind, ) -> ExecutionDecision: return ExecutionDecision( operator=operator, @@ -499,28 +519,63 @@ def _contract_problem_decision( strict: bool, reason: FallbackReason, ) -> ExecutionDecision: + decision: ExecutionDecisionKind = "strict-failure" if strict else "audit-warning" return ExecutionDecision( operator=operator, stage=stage, requested_mode="contract-check", requested_backend=None, actual_backend=None if strict else _NATIVE_BACKEND_ID, - decision="strict-failure" if strict else "audit-warning", + decision=decision, fallback=False, fallback_reason=reason, details=reason.details, ) +def execution_decision_sample_value( + *, + seed: int | str = 0, + rank: int | None = None, + key: str = "", +) -> float: + """Return a deterministic sample value in [0, 1) that is partitioned by rank.""" + rank = _current_process_rank() if rank is None else rank + payload = f"{seed}\0{rank}\0{key}".encode() + digest = hashlib.blake2b(payload, digest_size=8).digest() + return int.from_bytes(digest, "big") / 2**64 + + +def _current_process_rank() -> int: + for env_name in ("RANK", "LOCAL_RANK"): + value = os.environ.get(env_name) + if value is None: + continue + try: + return int(value) + except ValueError: + logger.debug("Ignoring non-integer %s=%r while sampling RL-Kernel decision logs.", env_name, value) + return 0 + + def emit_execution_decision( decision: ExecutionDecision, *, log: logging.Logger | None = None, enabled: bool = True, sample_rate: float = 1.0, - random_value: float = 0.0, + random_value: float | None = None, + sample_seed: int | str = 0, + sample_rank: int | None = None, + sample_key: str | None = None, ) -> dict[str, Any]: record = decision.to_log_record() + if random_value is None: + random_value = execution_decision_sample_value( + seed=sample_seed, + rank=sample_rank, + key=sample_key or json.dumps(record, sort_keys=True, default=str), + ) if not enabled or sample_rate <= 0 or random_value >= sample_rate: return record From c097537343a3bd5d219a1878c37b25428448133d Mon Sep 17 00:00:00 2001 From: inaniloquentee <3051000145@qq.com> Date: Thu, 23 Jul 2026 21:14:36 +0800 Subject: [PATCH 5/5] Productionize RL-Kernel linear_logp telemetry --- docs/en/advanced/rl-kernel-linear-logp.md | 111 +++ docs/en/index.rst | 1 + tests/_unit_stubs.py | 5 + .../test_rl_kernel_linear_logp_integration.py | 525 ++++++++++++++ vime/backends/megatron_utils/loss.py | 151 +++- vime/backends/megatron_utils/model.py | 96 ++- vime/backends/megatron_utils/rl_kernel.py | 674 ++++++++++++++++++ vime/backends/rl_kernel_utils/adapter.py | 35 +- vime/utils/rl_kernel.py | 51 ++ 9 files changed, 1618 insertions(+), 31 deletions(-) create mode 100644 docs/en/advanced/rl-kernel-linear-logp.md create mode 100644 tests/test_rl_kernel_linear_logp_integration.py create mode 100644 vime/backends/megatron_utils/rl_kernel.py create mode 100644 vime/utils/rl_kernel.py diff --git a/docs/en/advanced/rl-kernel-linear-logp.md b/docs/en/advanced/rl-kernel-linear-logp.md new file mode 100644 index 00000000..cdc80839 --- /dev/null +++ b/docs/en/advanced/rl-kernel-linear-logp.md @@ -0,0 +1,111 @@ +# RL-Kernel `linear_logp` + +vime can optionally use RL-Kernel for the actor selected-logprob path. When +`linear_logp` is enabled, the final Megatron pipeline stage returns hidden +states and vime calls the RL-Kernel operator adapter with hidden states, +LM-head weights, target token IDs, and tensor-parallel metadata. + +If the operator is disabled, unavailable, or unsupported for the current +runtime shape, vime falls back to the native Megatron output-layer plus +selected-logprob path. + +## Controls + +Preferred Phase 1 controls: + +```bash +--rlk-fast auto +--rl-kernel-ops linear_logp +``` + +Use strict mode when a missing or unsupported RL-Kernel path should fail the +run instead of falling back: + +```bash +--rlk-fast strict +--rl-kernel-ops linear_logp +``` + +Legacy aliases are still accepted and resolve into the same mode config: + +```bash +--enable-rl-kernel +--rl-kernel-strict +VIME_RLK_FAST=auto|strict +VIME_RL_KERNEL=1 +VIME_RL_KERNEL_STRICT=1 +VIME_RL_KERNEL_OPS=linear_logp +``` + +`VIME_LINEAR_LOGP_MEMORY_PROBE=1` enables optional CUDA memory probes around +the operator call. + +## Support Matrix + +| Backend | Implementation | dtype | Hardware/backend | TP | CP | Entropy | Full-gradient | +|---|---|---|---|---|---|---|---| +| `cuda_sm90` | RL-Kernel registry op, commonly `FusedLinearLogpSM90Op` when installed | Backend-defined; intended bf16/fp32 selected-logprob contracts are reported by RL-Kernel | NVIDIA SM90/Hopper CUDA backend when provided by the installed RL-Kernel package | Supported when the selected op accepts `tp_group`, `vocab_start_index`, and `global_vocab_size` | Not supported; falls back before CP redistribution | Not supported; falls back when entropy is requested | Supported only when the selected op saves hidden/weight backward state | +| `triton` | RL-Kernel registry op, commonly `TritonLinearLogpOp` when installed | Backend-defined floating input/output contract reported by RL-Kernel | CUDA devices supported by the installed RL-Kernel Triton backend | Supported when the selected op accepts TP metadata | Not supported; falls back before CP redistribution | Not supported; falls back when entropy is requested | Backend-defined; strict/full-gradient runs should validate saved-state support | +| `registry` | `RlkRegistryOperatorAdapter.linear_logp` calls `kernel_registry.get_op("linear_logp")` | Reported by the selected RL-Kernel backend; vime returns fp32 selected logprobs | Reported by the installed RL-Kernel backend | Supported when the selected op accepts TP metadata | Not supported; falls back before CP redistribution | Not supported; falls back when entropy is requested | Supported when the selected op saves hidden/weight backward state | +| `native` | Megatron output layer plus vime `calculate_log_probs_and_entropy` | vime native logits path, fp32 logprob computation | Same as native vime/Megatron execution | Supported by native vime/Megatron logprob path | Supported by native vime/Megatron CP redistribution path | Supported by native vime/Megatron path | Supported by native autograd over materialized logits | + +## Runtime Metadata + +Actor train-step logs include numeric metadata that is safe for W&B and +TensorBoard: + +```text +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_call_count_total +train/rl_kernel_linear_logp_call_count_delta +train/rl_kernel_linear_logp_token_count_total +train/rl_kernel_linear_logp_token_count_delta +train/rl_kernel_linear_logp_dispatch_elapsed_s_total +train/rl_kernel_linear_logp_dispatch_elapsed_s_delta +train/rl_kernel_linear_logp_tokens_per_call_total +train/rl_kernel_linear_logp_tokens_per_call_delta +``` + +The process log also prints human-readable fields: + +```text +requested_backend +actual_backend +backend_id +contract_id +fallback +fallback_reason +memory_probe_enabled +``` + +With `VIME_LINEAR_LOGP_MEMORY_PROBE=1`, vime adds operator-window memory +metrics when CUDA memory APIs are available: + +```text +train/rl_kernel_linear_logp_memory_alloc_delta_mb +train/rl_kernel_linear_logp_memory_peak_alloc_delta_mb +train/rl_kernel_linear_logp_memory_reserved_delta_mb +train/rl_kernel_linear_logp_memory_peak_reserved_delta_mb +``` + +These timing and memory values cover the `linear_logp` operator call only. Do +not report them as full-step speed or memory claims; full-step measurements +must include rollout, communication, optimizer, weight sync, and host +scheduling. + +## Fallback Behavior + +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. + +Common fallback reasons include: + +- the optional RL-Kernel package is unavailable; +- entropy was requested; +- CP redistribution is active; +- the selected op does not accept tensor-parallel metadata; +- the model output layer or LM-head weight is unavailable. diff --git a/docs/en/index.rst b/docs/en/index.rst index 7825277b..15b834e8 100644 --- a/docs/en/index.rst +++ b/docs/en/index.rst @@ -67,6 +67,7 @@ Start by Use Case advanced/vllm-config.md advanced/megatron-config.md advanced/rl-kernel-operator-adapter.md + advanced/rl-kernel-linear-logp.md advanced/arch-support-beyond-megatron.md .. toctree:: diff --git a/tests/_unit_stubs.py b/tests/_unit_stubs.py index afc27ec5..475db552 100644 --- a/tests/_unit_stubs.py +++ b/tests/_unit_stubs.py @@ -202,14 +202,18 @@ def install_megatron_mpu_stub() -> MagicMock: transformer_mod.__path__ = [] transformer_layer_mod = types.ModuleType("megatron.core.transformer.transformer_layer") transformer_layer_mod.get_transformer_layer_offset = lambda *args, **kwargs: 0 + tensor_parallel_mod = types.ModuleType("megatron.core.tensor_parallel") + tensor_parallel_mod.gather_from_sequence_parallel_region = lambda value, tensor_parallel_output_grad=False: value transformer_mod.transformer_layer = transformer_layer_mod megatron_core.parallel_state = parallel_state_mod megatron_core.transformer = transformer_mod + megatron_core.tensor_parallel = tensor_parallel_mod megatron_mod = types.ModuleType("megatron") megatron_mod.core = megatron_core sys.modules.setdefault("megatron", megatron_mod) sys.modules.setdefault("megatron.core", megatron_core) sys.modules.setdefault("megatron.core.parallel_state", parallel_state_mod) + sys.modules.setdefault("megatron.core.tensor_parallel", tensor_parallel_mod) sys.modules.setdefault("megatron.core.transformer", transformer_mod) sys.modules.setdefault("megatron.core.transformer.transformer_layer", transformer_layer_mod) return mpu_stub @@ -328,5 +332,6 @@ def install_triton_stub() -> None: def install_vime_distributed_utils_stub() -> None: vime_utils = types.ModuleType("vime.utils.distributed_utils") + vime_utils.distributed_masked_whiten = lambda values, masks, process_group=None, shift_mean=True: values vime_utils.get_gloo_group = MagicMock(return_value="gloo") sys.modules.setdefault("vime.utils.distributed_utils", vime_utils) diff --git a/tests/test_rl_kernel_linear_logp_integration.py b/tests/test_rl_kernel_linear_logp_integration.py new file mode 100644 index 00000000..0fe92b75 --- /dev/null +++ b/tests/test_rl_kernel_linear_logp_integration.py @@ -0,0 +1,525 @@ +from __future__ import annotations + +import importlib +import sys +import types +from argparse import Namespace +from pathlib import Path + +import pytest +import torch +import torch.nn.functional as F + +_tests_root = Path(__file__).resolve().parent +if str(_tests_root) not in sys.path: + sys.path.insert(0, str(_tests_root)) + +import _unit_stubs + +_unit_stubs.install_megatron_mpu_stub() +_unit_stubs.install_vime_distributed_utils_stub() + +from megatron.core import mpu # noqa: E402 + +from vime.backends.megatron_utils import loss as loss_mod # noqa: E402 +from vime.backends.megatron_utils import rl_kernel as rlk_mod # noqa: E402 + +adapter_mod = importlib.import_module("vime.backends.rl_kernel_utils.adapter") + + +class _FakeLinearLogpOp: + backend_id = "rlk.linear_logp.fake" + contract_id = "rlk.linear_logp.fake.fp32" + calls: list[dict] = [] + + def __call__( + self, + hidden: torch.Tensor, + weight: torch.Tensor, + target_ids: torch.Tensor, + bias: torch.Tensor | None = None, + **kwargs, + ) -> torch.Tensor: + type(self).calls.append( + { + "hidden_shape": tuple(hidden.shape), + "weight_shape": tuple(weight.shape), + "target_shape": tuple(target_ids.shape), + "bias": bias is not None, + "kwargs": kwargs, + "hidden_requires_grad": hidden.requires_grad, + "hidden_dtype": hidden.dtype, + } + ) + logits = F.linear(hidden.float(), weight.float(), None if bias is None else bias.float()) + return torch.gather(torch.log_softmax(logits, dim=-1), -1, target_ids.long().unsqueeze(-1)).squeeze(-1) + + +class _FakeLegacyLinearLogpOp: + def __call__( + self, + hidden: torch.Tensor, + weight: torch.Tensor, + target_ids: torch.Tensor, + bias: torch.Tensor | None = None, + ) -> torch.Tensor: + logits = F.linear(hidden.float(), weight.float(), None if bias is None else bias.float()) + return torch.gather(torch.log_softmax(logits, dim=-1), -1, target_ids.long().unsqueeze(-1)).squeeze(-1) + + +def _drop_rl_engine_modules() -> None: + for name in list(sys.modules): + if name == "rl_engine" or name.startswith("rl_engine."): + sys.modules.pop(name, None) + + +def _reset_rl_kernel_state() -> None: + rlk_mod._LINEAR_LOGP_ADAPTER = None + rlk_mod._LINEAR_LOGP_ADAPTER_ERROR = None + rlk_mod._LINEAR_LOGP_SAVE_PROBS_CAST_LOGGED = False + rlk_mod._WARNED_FALLBACK_REASONS.clear() + rlk_mod._FALLBACK_COUNTS.clear() + rlk_mod._FALLBACK_COUNTS.update({"linear_logp": 0}) + rlk_mod.reset_rl_kernel_runtime_counters() + _FakeLinearLogpOp.calls.clear() + _drop_rl_engine_modules() + + +def _install_fake_rl_engine(monkeypatch, op_factory=lambda: _FakeLinearLogpOp()) -> None: + rl_engine = types.ModuleType("rl_engine") + kernels = types.ModuleType("rl_engine.kernels") + registry = types.ModuleType("rl_engine.kernels.registry") + registry.kernel_registry = types.SimpleNamespace(get_op=lambda name: op_factory()) + monkeypatch.setitem(sys.modules, "rl_engine", rl_engine) + monkeypatch.setitem(sys.modules, "rl_engine.kernels", kernels) + monkeypatch.setitem(sys.modules, "rl_engine.kernels.registry", registry) + + +def _make_args(**overrides) -> Namespace: + values = { + "rlk_fast": "auto", + "rlk_consistency": "off", + "rlk_mode_config": types.SimpleNamespace(fast="auto", consistency="off", ops=("linear_logp",)), + "enable_rl_kernel": True, + "rl_kernel_ops": ("linear_logp",), + "rl_kernel_strict": False, + "allgather_cp": False, + "qkv_format": "thd", + "rollout_temperature": 1.0, + "log_probs_chunk_size": -1, + "entropy_coef": 0.0, + "sequence_parallel": False, + "padded_vocab_size": None, + "only_train_params_name_list": (), + } + values.update(overrides) + if "rlk_mode_config" not in overrides: + values["rlk_mode_config"] = types.SimpleNamespace( + fast=values["rlk_fast"], + consistency=values["rlk_consistency"], + ops=values["rl_kernel_ops"], + ) + return Namespace(**values) + + +@pytest.fixture(autouse=True) +def reset_parallelism(): + _reset_rl_kernel_state() + mpu.get_tensor_model_parallel_world_size.return_value = 1 + mpu.get_tensor_model_parallel_rank.return_value = 0 + mpu.get_tensor_model_parallel_group.return_value = None + mpu.get_context_parallel_world_size.return_value = 1 + mpu.get_context_parallel_rank.return_value = 0 + mpu.get_virtual_pipeline_model_parallel_world_size.return_value = None + mpu.is_pipeline_last_stage.return_value = True + yield + _reset_rl_kernel_state() + + +def _reference_logp( + hidden: torch.Tensor, + weight: torch.Tensor, + target: torch.Tensor, + bias: torch.Tensor | None, +) -> torch.Tensor: + logits = F.linear(hidden.float(), weight.float(), None if bias is None else bias.float()) + return torch.gather(torch.log_softmax(logits, dim=-1), -1, target.long().unsqueeze(-1)).squeeze(-1) + + +def _cpu_calculate_log_probs_and_entropy( + logits: torch.Tensor, + tokens: torch.Tensor, + tp_group, + *, + with_entropy: bool, + chunk_size: int, + log_prob_keep_mask=None, + with_entropy_grad: bool = True, +): + del tp_group, chunk_size, with_entropy_grad + masked_logits = logits.float() + if log_prob_keep_mask is not None: + masked_logits = masked_logits.masked_fill(~log_prob_keep_mask, float("-inf")) + log_probs = torch.log_softmax(masked_logits, dim=-1) + selected = torch.gather(log_probs, -1, tokens.long().unsqueeze(-1)).squeeze(-1) + entropy = None + if with_entropy: + native_log_probs = torch.log_softmax(logits.float(), dim=-1) + probs = torch.softmax(logits.float(), dim=-1) + entropy = -(probs * native_log_probs).sum(dim=-1) + return selected, entropy + + +@pytest.mark.unit +def test_maybe_compute_linear_logp_passes_tensor_parallel_metadata(monkeypatch): + _install_fake_rl_engine(monkeypatch) + args = _make_args() + torch.manual_seed(1) + hidden = torch.randn(6, 5) + weight = torch.randn(8, 5) + bias = torch.randn(8) + target = torch.randint(0, 8, (6,)) + context = rlk_mod.LinearLogpContext( + lm_head_weight=weight, + bias=bias, + tp_group="tp", + vocab_start_index=16, + global_vocab_size=32, + ) + + actual = rlk_mod.maybe_compute_linear_logp(hidden, target, context=context, args=args, with_entropy=False) + + torch.testing.assert_close(actual, _reference_logp(hidden, weight, target, bias)) + assert _FakeLinearLogpOp.calls == [ + { + "hidden_shape": (6, 5), + "weight_shape": (8, 5), + "target_shape": (6,), + "bias": True, + "kwargs": { + "tp_group": "tp", + "vocab_start_index": 16, + "global_vocab_size": 32, + }, + "hidden_requires_grad": False, + "hidden_dtype": hidden.dtype, + } + ] + counters = rlk_mod.get_rl_kernel_runtime_counters() + assert counters["linear_logp_call_count"] == 1.0 + assert counters["linear_logp_token_count"] == 6.0 + assert counters["linear_logp_dispatch_elapsed_s"] >= 0.0 + assert counters["linear_logp_fallback_count"] == 0.0 + + metadata = rlk_mod.get_linear_logp_runtime_metadata() + assert metadata["requested_backend"] == "registry" + assert metadata["actual_backend"] == "_FakeLinearLogpOp" + assert metadata["backend_id"] == "rlk.linear_logp.fake" + assert metadata["contract_id"] == "rlk.linear_logp.fake.fp32" + assert metadata["fallback"] is False + assert metadata["fallback_reason"] is None + + log_metrics = rlk_mod.get_linear_logp_runtime_log_metrics(prefix="x/") + assert log_metrics["x/fallback"] == 0.0 + assert log_metrics["x/backend_descriptor_id"] == metadata["backend_descriptor_id"] + assert log_metrics["x/contract_descriptor_id"] == metadata["contract_descriptor_id"] + assert log_metrics["x/fallback_reason_descriptor_id"] == metadata["fallback_reason_descriptor_id"] + assert all("memory_" not in key for key in log_metrics) + + +@pytest.mark.unit +def test_linear_logp_full_gradient_path_matches_materialized_logits(monkeypatch): + _install_fake_rl_engine(monkeypatch) + args = _make_args() + torch.manual_seed(11) + hidden = torch.randn(5, 4, requires_grad=True) + weight = torch.randn(7, 4, requires_grad=True) + bias = torch.randn(7, requires_grad=True) + target = torch.randint(0, 7, (5,)) + context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=bias, tp_group=None) + + actual = rlk_mod.maybe_compute_linear_logp(hidden, target, context=context, args=args, with_entropy=False) + actual.sum().backward() + actual_grads = (hidden.grad.clone(), weight.grad.clone(), bias.grad.clone()) + + hidden_ref = hidden.detach().clone().requires_grad_(True) + weight_ref = weight.detach().clone().requires_grad_(True) + bias_ref = bias.detach().clone().requires_grad_(True) + expected = _reference_logp(hidden_ref, weight_ref, target, bias_ref) + expected.sum().backward() + + torch.testing.assert_close(actual.detach(), expected.detach(), rtol=1e-6, atol=1e-6) + torch.testing.assert_close(actual_grads[0], hidden_ref.grad, rtol=1e-6, atol=1e-6) + torch.testing.assert_close(actual_grads[1], weight_ref.grad, rtol=1e-6, atol=1e-6) + torch.testing.assert_close(actual_grads[2], bias_ref.grad, rtol=1e-6, atol=1e-6) + + +@pytest.mark.unit +def test_linear_logp_matches_vime_response_slicing_from_hidden_states(monkeypatch): + _install_fake_rl_engine(monkeypatch) + args = _make_args() + vocab_size = 23 + hidden_size = 7 + total_lengths = [5, 4, 6] + response_lengths = [2, 3, 4] + torch.manual_seed(2) + unconcat_tokens = [torch.randint(0, vocab_size, (length,), dtype=torch.long) for length in total_lengths] + hidden = torch.randn(sum(total_lengths), 1, hidden_size) + weight = torch.randn(vocab_size, hidden_size) + bias = torch.randn(vocab_size) + context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=bias, tp_group=None) + + _, result = loss_mod.get_log_probs_and_entropy( + hidden, + args=args, + unconcat_tokens=unconcat_tokens, + total_lengths=total_lengths, + response_lengths=response_lengths, + with_entropy=False, + rl_kernel_linear_logp_context=context, + ) + + full_tokens = torch.zeros(sum(total_lengths), dtype=torch.long) + offset = 0 + for tokens, total_length in zip(unconcat_tokens, total_lengths, strict=False): + full_tokens[offset : offset + total_length - 1] = tokens[1:total_length] + offset += total_length + full_logp = _reference_logp(hidden.squeeze(1), weight, full_tokens, bias) + + expected = [] + offset = 0 + for total_length, response_length in zip(total_lengths, response_lengths, strict=False): + end = offset + total_length + start = end - response_length + expected.append(full_logp[start - 1 : end - 1]) + offset += total_length + + assert len(_FakeLinearLogpOp.calls) == 1 + for actual_item, expected_item in zip(result["log_probs"], expected, strict=True): + torch.testing.assert_close(actual_item, expected_item, rtol=1e-6, atol=1e-6) + + +@pytest.mark.unit +def test_linear_logp_materializes_logits_fallback_when_optional_package_missing(monkeypatch): + monkeypatch.setattr(loss_mod, "calculate_log_probs_and_entropy", _cpu_calculate_log_probs_and_entropy) + original_import_module = adapter_mod.importlib.import_module + + def fail_rl_engine_import(name, *args, **kwargs): + if name == "rl_engine.kernels.registry": + raise ModuleNotFoundError("No module named 'rl_engine'") + return original_import_module(name, *args, **kwargs) + + monkeypatch.setattr(adapter_mod.importlib, "import_module", fail_rl_engine_import) + args = _make_args() + vocab_size = 19 + hidden_size = 5 + total_lengths = [4, 5] + response_lengths = [2, 3] + torch.manual_seed(3) + unconcat_tokens = [torch.randint(0, vocab_size, (length,), dtype=torch.long) for length in total_lengths] + hidden = torch.randn(1, sum(total_lengths), hidden_size) + weight = torch.randn(vocab_size, hidden_size) + bias = torch.randn(vocab_size) + context = rlk_mod.LinearLogpContext(lm_head_weight=weight, bias=bias, tp_group=None) + + _, result = loss_mod.get_log_probs_and_entropy( + hidden, + args=args, + unconcat_tokens=unconcat_tokens, + total_lengths=total_lengths, + response_lengths=response_lengths, + with_entropy=False, + rl_kernel_linear_logp_context=context, + ) + + logits = F.linear(hidden.squeeze(0).float(), weight.float(), bias.float()) + expected = [] + offset = 0 + for tokens, total_length, response_length in zip(unconcat_tokens, total_lengths, response_lengths, strict=False): + end = offset + total_length + start = end - response_length + target = tokens[-response_length:] + expected.append(torch.gather(torch.log_softmax(logits[start - 1 : end - 1], dim=-1), -1, target.unsqueeze(-1)).squeeze(-1)) + offset += total_length + + 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 "No module named 'rl_engine'" in metadata["fallback_reason"] + log_metrics = rlk_mod.get_linear_logp_runtime_log_metrics(prefix="x/") + assert log_metrics["x/fallback"] == 1.0 + assert log_metrics["x/backend_descriptor_id"] == metadata["backend_descriptor_id"] + assert log_metrics["x/contract_descriptor_id"] == metadata["contract_descriptor_id"] + assert log_metrics["x/fallback_reason_descriptor_id"] == metadata["fallback_reason_descriptor_id"] + for actual_item, expected_item in zip(result["log_probs"], expected, strict=True): + torch.testing.assert_close(actual_item, expected_item, rtol=1e-6, atol=1e-6) + + +@pytest.mark.unit +def test_linear_logp_falls_back_when_op_lacks_tp_interface(monkeypatch): + _install_fake_rl_engine(monkeypatch, op_factory=lambda: _FakeLegacyLinearLogpOp()) + args = _make_args() + 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="tp_group", + vocab_start_index=6, + global_vocab_size=12, + ) + + actual = rlk_mod.maybe_compute_linear_logp(hidden, target, context=context, args=args, with_entropy=False) + + assert actual is None + assert rlk_mod.get_rl_kernel_fallback_count("linear_logp") == 1 + metadata = rlk_mod.get_linear_logp_runtime_metadata() + assert metadata["fallback"] is True + assert metadata["actual_backend"] == "vime.native.linear_logp" + assert "unexpected keyword argument" in metadata["fallback_reason"] + + +@pytest.mark.unit +def test_linear_logp_zero_tokens_reports_non_fallback_decision(monkeypatch): + _install_fake_rl_engine(monkeypatch) + args = _make_args() + hidden = torch.randn(0, 3) + weight = torch.randn(5, 3) + target = torch.empty(0, dtype=torch.long) + 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.shape == (0,) + counters = rlk_mod.get_rl_kernel_runtime_counters() + assert counters["linear_logp_call_count"] == 0.0 + assert counters["linear_logp_fallback_count"] == 0.0 + metadata = rlk_mod.get_linear_logp_runtime_metadata() + assert metadata["actual_backend"] == "vime.linear_logp.zero_tokens" + assert metadata["fallback"] is False + assert metadata["fallback_reason"] is None + + +@pytest.mark.unit +def test_linear_logp_support_matrix_documents_issue20_fields(): + matrix = rlk_mod.get_linear_logp_support_matrix() + assert matrix + + required_fields = { + "backend", + "implementation", + "dtype", + "hardware", + "tp", + "cp", + "entropy", + "full_gradient", + } + assert all(required_fields <= set(row) for row in matrix) + assert {row["backend"] for row in matrix} >= {"registry", "native"} + assert any(row["source"] == "vime_adapter" for row in matrix) + assert all(row["source"] == "rl_kernel" for row in matrix if row["backend"] in {"cuda_sm90", "triton", "pytorch"}) + assert any("native" in row["backend"] and "Megatron" in row["implementation"] for row in matrix) + + +@pytest.mark.unit +def test_linear_logp_support_matrix_reports_unavailable_provider(monkeypatch): + def fail_import(name, *args, **kwargs): + if name == "rl_engine.kernels.support": + raise ModuleNotFoundError("No module named 'rl_engine'") + return importlib.import_module(name, *args, **kwargs) + + monkeypatch.setattr(rlk_mod.importlib, "import_module", fail_import) + + matrix = rlk_mod.get_linear_logp_support_matrix() + + assert {row["backend"] for row in matrix} >= {"rl_kernel_unavailable", "registry", "native"} + + +@pytest.mark.unit +def test_linear_logp_support_matrix_imports_rl_kernel_rows(monkeypatch): + _drop_rl_engine_modules() + rl_engine = types.ModuleType("rl_engine") + kernels = types.ModuleType("rl_engine.kernels") + support = types.ModuleType("rl_engine.kernels.support") + support.get_linear_logp_support_matrix = lambda: ( + { + "source": "rl_kernel", + "backend": "cuda_sm90", + "implementation": "FusedLinearLogpSM90Op", + "dtype": "backend-defined", + "hardware": "SM90", + "tp": "reported by RL-Kernel", + "cp": "reported by RL-Kernel", + "entropy": "reported by RL-Kernel", + "full_gradient": "reported by RL-Kernel", + }, + ) + monkeypatch.setitem(sys.modules, "rl_engine", rl_engine) + monkeypatch.setitem(sys.modules, "rl_engine.kernels", kernels) + monkeypatch.setitem(sys.modules, "rl_engine.kernels.support", support) + + matrix = rlk_mod.get_linear_logp_support_matrix() + + assert {row["backend"] for row in matrix} >= {"cuda_sm90", "registry", "native"} + assert matrix[0]["source"] == "rl_kernel" + + +@pytest.mark.unit +def test_linear_logp_context_from_model_uses_tp_vocab_offsets(): + mpu.get_tensor_model_parallel_world_size.return_value = 4 + mpu.get_tensor_model_parallel_rank.return_value = 2 + mpu.get_tensor_model_parallel_group.return_value = "tp_group" + + output_layer = types.SimpleNamespace( + weight=torch.empty(8, 4), + bias=None, + sequence_parallel=True, + ) + model = types.SimpleNamespace(output_layer=output_layer, post_process=True) + args = _make_args(padded_vocab_size=32, sequence_parallel=False) + + context = rlk_mod.get_linear_logp_context_from_model(args, model) + + assert context is not None + assert context.lm_head_weight is output_layer.weight + assert context.tp_group == "tp_group" + assert context.vocab_start_index == 16 + assert context.global_vocab_size == 32 + assert context.sequence_parallel is True + + +@pytest.mark.unit +def test_return_hidden_states_for_linear_logp_restores_post_process_flag(): + args = _make_args() + model = types.SimpleNamespace(post_process=True) + context = rlk_mod.LinearLogpContext( + lm_head_weight=torch.empty(4, 3), + bias=None, + tp_group=None, + ) + + with rlk_mod.return_hidden_states_for_linear_logp(args, model, context) as enabled: + assert enabled is True + assert model.post_process is False + + assert model.post_process is True + + +@pytest.mark.unit +def test_policy_loss_only_skips_entropy_when_linear_logp_context_is_active(): + args = _make_args(entropy_coef=0.0) + context = rlk_mod.LinearLogpContext( + lm_head_weight=torch.empty(4, 3), + bias=None, + tp_group=None, + ) + + assert loss_mod._policy_loss_needs_entropy(args, None) is True + assert loss_mod._policy_loss_needs_entropy(args, context) is False + + args.entropy_coef = 0.01 + assert loss_mod._policy_loss_needs_entropy(args, None) is True + assert loss_mod._policy_loss_needs_entropy(args, context) is True diff --git a/vime/backends/megatron_utils/loss.py b/vime/backends/megatron_utils/loss.py index c24b5323..90f36b93 100644 --- a/vime/backends/megatron_utils/loss.py +++ b/vime/backends/megatron_utils/loss.py @@ -24,6 +24,7 @@ get_reinforce_plus_plus_baseline_advantages, get_reinforce_plus_plus_returns, ) +from vime.utils.rl_kernel import is_rl_kernel_op_enabled from vime.utils.types import RolloutBatch from .cp_utils import ( @@ -32,6 +33,12 @@ get_sum_of_sample_mean, slice_log_prob_with_cp, ) +from .rl_kernel import ( + LinearLogpContext, + get_rl_kernel_fallback_count, + maybe_compute_linear_logp, + warn_linear_logp_fallback, +) ROLLOUT_TOP_P_TOKEN_KEYS = ( "rollout_top_p_token_ids", @@ -481,6 +488,54 @@ def _extract_per_sample( return log_probs_list, entropy_list +def _gather_sequence_parallel_hidden_if_needed( + hidden_states: torch.Tensor, + context: LinearLogpContext | None, +) -> torch.Tensor: + if context is None or not context.sequence_parallel: + return hidden_states + + from megatron.core import tensor_parallel + + return tensor_parallel.gather_from_sequence_parallel_region(hidden_states, tensor_parallel_output_grad=False) + + +def _flatten_logprob_model_output( + output_tensor: torch.Tensor, + *, + linear_logp_context: LinearLogpContext | None, +) -> torch.Tensor: + assert len(output_tensor.shape) == 3, f"{output_tensor.shape}" + if output_tensor.size(0) == 1: + return output_tensor.squeeze(0) + if linear_logp_context is not None and output_tensor.size(1) == 1: + return output_tensor.squeeze(1) + assert output_tensor.size(0) == 1, f"{output_tensor.shape}" + return output_tensor.squeeze(0) + + +def _materialize_linear_logits( + hidden_states: torch.Tensor, + *, + context: LinearLogpContext, + args: Namespace, +) -> torch.Tensor: + logits = F.linear(hidden_states, context.lm_head_weight, context.bias) + rollout_temperature = getattr(args, "rollout_temperature", 1.0) + if rollout_temperature != 1.0: + logits = logits / rollout_temperature + return logits.float() + + +def _policy_loss_needs_entropy( + args: Namespace, + rl_kernel_linear_logp_context: LinearLogpContext | None, +) -> bool: + if rl_kernel_linear_logp_context is None: + return True + return getattr(args, "entropy_coef", 0.0) != 0 + + def get_log_probs_and_entropy( logits: torch.Tensor, *, @@ -492,6 +547,7 @@ def get_log_probs_and_entropy( non_loss_data: bool = True, top_p_token_ids: list[list[int]] | None = None, top_p_token_offsets: list[list[int]] | None = None, + rl_kernel_linear_logp_context: LinearLogpContext | None = None, ) -> dict[str, list[torch.Tensor]]: """Compute per-token log-probabilities (and optionally entropy) on responses. @@ -503,16 +559,19 @@ def get_log_probs_and_entropy( log-probabilities; entropy is always computed from the unmasked logits. """ assert non_loss_data - assert logits.dtype == torch.float32, f"{logits.dtype}" - assert len(logits.shape) == 3, f"{logits.shape}" - assert logits.size(0) == 1, f"{logits.shape}" - logits = logits.squeeze(0) - - # Apply rollout temperature scaling to logits to match rollout-time log-probs. - rollout_temperature = getattr(args, "rollout_temperature", 1.0) - if rollout_temperature != 1.0: - logits = logits / rollout_temperature - logits = logits.contiguous() + linear_logp_context = rl_kernel_linear_logp_context + if linear_logp_context is not None: + logits = _gather_sequence_parallel_hidden_if_needed(logits, linear_logp_context) + else: + assert logits.dtype == torch.float32, f"{logits.dtype}" + + logits = _flatten_logprob_model_output(logits, linear_logp_context=linear_logp_context).contiguous() + if linear_logp_context is None: + # Apply rollout temperature scaling to logits to match rollout-time log-probs. + rollout_temperature = getattr(args, "rollout_temperature", 1.0) + if rollout_temperature != 1.0: + logits = logits / rollout_temperature + logits = logits.contiguous() T = logits.size(0) device = logits.device tp_group = mpu.get_tensor_model_parallel_group() @@ -524,6 +583,22 @@ def get_log_probs_and_entropy( # --- build full shifted-token target tensor --- full_tokens = _build_shifted_tokens(T, device, unconcat_tokens, total_lengths, response_lengths, args.allgather_cp) + log_prob_full = None + if linear_logp_context is not None and (top_p_token_ids is not None or top_p_token_offsets is not None): + warn_linear_logp_fallback(args, "rollout top-p replay requires materialized logits") + elif linear_logp_context is not None: + log_prob_full = maybe_compute_linear_logp( + logits, + full_tokens, + context=linear_logp_context, + args=args, + with_entropy=with_entropy, + ) + + if log_prob_full is None and linear_logp_context is not None: + logits = _materialize_linear_logits(logits, context=linear_logp_context, args=args).contiguous() + linear_logp_context = None + # --- build top-p nucleus keep-mask (logprob only; entropy stays unmasked) --- top_p_keep_mask = None if top_p_token_ids is not None and top_p_token_offsets is not None: @@ -539,15 +614,18 @@ def get_log_probs_and_entropy( ) # --- compute on full [T,V] logits at once via calculate_log_probs_and_entropy --- - log_prob_full, entropy_full = calculate_log_probs_and_entropy( - logits, - full_tokens, - tp_group, - with_entropy=with_entropy, - with_entropy_grad=with_entropy_grad, - chunk_size=chunk_size, - log_prob_keep_mask=top_p_keep_mask, - ) + if log_prob_full is None: + log_prob_full, entropy_full = calculate_log_probs_and_entropy( + logits, + full_tokens, + tp_group, + with_entropy=with_entropy, + with_entropy_grad=with_entropy_grad, + chunk_size=chunk_size, + log_prob_keep_mask=top_p_keep_mask, + ) + else: + entropy_full = None log_prob_full = log_prob_full.squeeze(-1) # [T, 1] -> [T] # --- extract per-sample response portions --- @@ -897,6 +975,7 @@ def policy_loss_function( batch: RolloutBatch, logits: torch.Tensor, sum_of_sample_mean: Callable[[torch.Tensor], torch.Tensor], + rl_kernel_linear_logp_context: LinearLogpContext | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute policy loss (PPO/GSPO) and metrics. @@ -927,6 +1006,7 @@ def policy_loss_function( response_lengths = batch["response_lengths"] total_lengths = batch["total_lengths"] + need_entropy = _policy_loss_needs_entropy(args, rl_kernel_linear_logp_context) _, log_probs_and_entropy = get_log_probs_and_entropy( logits, @@ -934,7 +1014,8 @@ def policy_loss_function( unconcat_tokens=batch["unconcat_tokens"], total_lengths=total_lengths, response_lengths=response_lengths, - with_entropy=True, + with_entropy=need_entropy, + rl_kernel_linear_logp_context=rl_kernel_linear_logp_context, **get_rollout_top_p_logprob_kwargs(args, batch), ) @@ -1059,9 +1140,12 @@ def policy_loss_function( ppo_kl = sum_of_sample_mean(ppo_kl) # entropy loss - entropy = log_probs_and_entropy["entropy"] - entropy = torch.cat(entropy, dim=0) - entropy_loss = sum_of_sample_mean(entropy) + if need_entropy: + entropy = log_probs_and_entropy["entropy"] + entropy = torch.cat(entropy, dim=0) + entropy_loss = sum_of_sample_mean(entropy) + else: + entropy_loss = log_probs.new_zeros(()) loss = pg_loss - args.entropy_coef * entropy_loss @@ -1149,6 +1233,7 @@ def value_loss_function( batch: RolloutBatch, logits: torch.Tensor, sum_of_sample_mean: Callable[[torch.Tensor], torch.Tensor], + rl_kernel_linear_logp_context: LinearLogpContext | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute clipped value loss and metrics. @@ -1167,6 +1252,7 @@ def value_loss_function( Tuple of `(loss, metrics)` where `loss` is a scalar tensor and `metrics` contains detached scalars "value_loss" and "value_clipfrac". """ + del rl_kernel_linear_logp_context old_values = torch.cat(batch["values"], dim=0) _, values = get_values( @@ -1206,6 +1292,7 @@ def sft_loss_function( batch: RolloutBatch, logits: torch.Tensor, sum_of_sample_mean: Callable[[torch.Tensor], torch.Tensor], + rl_kernel_linear_logp_context: LinearLogpContext | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """Compute supervised fine-tuning loss over response tokens. @@ -1233,6 +1320,7 @@ def sft_loss_function( total_lengths=total_lengths, response_lengths=response_lengths, with_entropy=False, + rl_kernel_linear_logp_context=rl_kernel_linear_logp_context, ) log_probs = log_probs_and_entropy["log_probs"] @@ -1257,6 +1345,7 @@ def loss_function( num_microbatches: int, step_global_batch_size: int, logits: torch.Tensor, + rl_kernel_linear_logp_context: LinearLogpContext | None = None, ) -> tuple[torch.Tensor, int | torch.Tensor, dict[str, list[str] | torch.Tensor]]: """Dispatch to the configured loss and rescale for Megatron integration. @@ -1307,10 +1396,22 @@ def loss_function( case _: raise ValueError(f"Unknown loss type: {args.loss_type}") + if func in {policy_loss_function, value_loss_function, sft_loss_function}: + func_args = (args, batch, logits, sum_of_sample_mean, rl_kernel_linear_logp_context) + else: + func_args = (args, batch, logits, sum_of_sample_mean) + if args.recompute_loss_function: - loss, log = checkpoint(func, args, batch, logits, sum_of_sample_mean, use_reentrant=False) + loss, log = checkpoint(func, *func_args, use_reentrant=False) else: - loss, log = func(args, batch, logits, sum_of_sample_mean) + loss, log = func(*func_args) + + if is_rl_kernel_op_enabled(args, "linear_logp"): + log["rl_kernel_fallback_count"] = torch.tensor( + get_rl_kernel_fallback_count(), + device=logits.device, + dtype=torch.float32, + ) # With allgather-CP, some CP ranks may have no loss-contributing tokens (e.g., all # padding). Without this, gradient doesn't flow through their attention path, so diff --git a/vime/backends/megatron_utils/model.py b/vime/backends/megatron_utils/model.py index 8642927a..5891b736 100644 --- a/vime/backends/megatron_utils/model.py +++ b/vime/backends/megatron_utils/model.py @@ -30,12 +30,23 @@ 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 from .checkpoint import load_checkpoint, save_checkpoint from .cp_utils import reduce_train_step_metrics from .data import DataIterator, get_batch -from .loss import ROLLOUT_TOP_P_TOKEN_KEYS, get_rollout_top_p_logprob_kwargs, loss_function +from .loss import ROLLOUT_TOP_P_TOKEN_KEYS, get_log_probs_and_entropy, get_rollout_top_p_logprob_kwargs, loss_function 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_rl_kernel_runtime_counter_delta, + get_rl_kernel_runtime_counters, + return_hidden_states_for_linear_logp, + should_use_linear_logp_model_output, + warn_linear_logp_fallback, +) from .stateless_adam import StatelessAdam logger = logging.getLogger(__name__) @@ -81,6 +92,35 @@ def _with_rollout_top_p_token_keys(args: Namespace, keys: Sequence[str]) -> list return [*keys, *ROLLOUT_TOP_P_TOKEN_KEYS] +def _forward_only_should_return_hidden_for_linear_logp( + f: Callable[..., dict[str, list[torch.Tensor]]], + args: Namespace, +) -> bool: + return f is get_log_probs_and_entropy and should_use_linear_logp_model_output( + args, + with_entropy=args.use_rollout_entropy, + ) + + +def _train_should_return_hidden_for_linear_logp(args: Namespace, *, return_schedule_plan: bool) -> bool: + if not is_rl_kernel_op_enabled(args, "linear_logp"): + return False + + if args.loss_type not in {"policy_loss", "sft_loss"}: + return False + + if return_schedule_plan: + warn_linear_logp_fallback(args, "schedule-plan forward path is not supported") + return False + + if getattr(args, "enable_mtp_training", False): + warn_linear_logp_fallback(args, "MTP training path is not supported") + return False + + with_entropy = args.loss_type == "policy_loss" and getattr(args, "entropy_coef", 0.0) != 0 + return should_use_linear_logp_model_output(args, with_entropy=with_entropy) + + def _iter_critic_output_layers(model: Sequence[DDP]): for chunk_id, module in enumerate(unwrap_model(model)): output_layer = getattr(module, "output_layer", None) @@ -430,7 +470,12 @@ def forward_step( } if batch["multimodal_train_inputs"] is not None: forward_kwargs.update(batch["multimodal_train_inputs"]) - output_tensor = model(**forward_kwargs) + linear_logp_context = None + if _forward_only_should_return_hidden_for_linear_logp(f, args): + linear_logp_context = get_linear_logp_context_from_model(args, model) + + with return_hidden_states_for_linear_logp(args, model, linear_logp_context): + output_tensor = model(**forward_kwargs) output_kwargs = { "args": args, @@ -441,6 +486,8 @@ def forward_step( } if use_rollout_top_p_replay: output_kwargs.update(get_rollout_top_p_logprob_kwargs(args, batch)) + if f is get_log_probs_and_entropy: + output_kwargs["rl_kernel_linear_logp_context"] = linear_logp_context return output_tensor, partial(f, **output_kwargs) @@ -603,6 +650,10 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p old_stage = os.environ["ROUTING_REPLAY_STAGE"] os.environ["ROUTING_REPLAY_STAGE"] = "replay_forward" + linear_logp_context = None + if _train_should_return_hidden_for_linear_logp(args, return_schedule_plan=return_schedule_plan): + linear_logp_context = get_linear_logp_context_from_model(args, model) + if return_schedule_plan: assert not args.enable_mtp_training, "MTP training should not be enabled when using combined 1f1b" position_ids = None @@ -646,12 +697,20 @@ def forward_step(data_iterator: DataIterator, model: GPTModel, return_schedule_p if args.enable_mtp_training: forward_kwargs["mtp_kwargs"] = {"mtp_labels": batch["tokens"]} - output_tensor = model(**forward_kwargs) + with return_hidden_states_for_linear_logp(args, model, linear_logp_context): + output_tensor = model(**forward_kwargs) if os.environ.get("ENABLE_ROUTING_REPLAY", "0") == "1": os.environ["ROUTING_REPLAY_STAGE"] = old_stage - return output_tensor, partial(loss_function, args, batch, num_microbatches, step_global_batch_size) + return output_tensor, partial( + loss_function, + args, + batch, + num_microbatches, + step_global_batch_size, + rl_kernel_linear_logp_context=linear_logp_context, + ) # Forward pass. forward_backward_func = get_forward_backward_func() @@ -902,6 +961,35 @@ def train( # Per-step gbs — uneven step sizes are easy to miss without this. log_dict[f"train/{role_tag}global_batch_size"] = global_batch_sizes[step_id] + if role == "actor" and is_rl_kernel_op_enabled(args, "linear_logp"): + runtime_totals = get_rl_kernel_runtime_counters() + runtime_delta = get_rl_kernel_runtime_counter_delta() + linear_logp_metadata = get_linear_logp_runtime_metadata() + for key, value in runtime_totals.items(): + log_dict[f"train/rl_kernel_{key}_total"] = value + for key, value in runtime_delta.items(): + log_dict[f"train/rl_kernel_{key}_delta"] = value + log_dict.update(get_linear_logp_runtime_log_metrics()) + + total_calls = runtime_totals.get("linear_logp_call_count", 0.0) + delta_calls = runtime_delta.get("linear_logp_call_count", 0.0) + log_dict["train/rl_kernel_linear_logp_tokens_per_call_total"] = ( + runtime_totals.get("linear_logp_token_count", 0.0) / total_calls if total_calls > 0 else 0.0 + ) + log_dict["train/rl_kernel_linear_logp_tokens_per_call_delta"] = ( + runtime_delta.get("linear_logp_token_count", 0.0) / delta_calls if delta_calls > 0 else 0.0 + ) + logger.info( + "RL-Kernel linear_logp runtime_metadata: requested_backend=%s actual_backend=%s " + "backend_id=%s contract_id=%s fallback=%s fallback_reason=%s memory_probe_enabled=%s", + linear_logp_metadata.get("requested_backend"), + linear_logp_metadata.get("actual_backend"), + linear_logp_metadata.get("backend_id"), + linear_logp_metadata.get("contract_id"), + linear_logp_metadata.get("fallback"), + linear_logp_metadata.get("fallback_reason"), + linear_logp_metadata.get("memory_probe_enabled"), + ) log_dict["train/step"] = accumulated_step_id logging_utils.log(args, log_dict, step_key="train/step") diff --git a/vime/backends/megatron_utils/rl_kernel.py b/vime/backends/megatron_utils/rl_kernel.py new file mode 100644 index 00000000..71f79dea --- /dev/null +++ b/vime/backends/megatron_utils/rl_kernel.py @@ -0,0 +1,674 @@ +from __future__ import annotations + +import hashlib +import importlib +import logging +import os +import re +import time +from argparse import Namespace +from contextlib import contextmanager +from dataclasses import asdict, dataclass +from typing import Any + +import torch +from megatron.core import mpu + +from vime.backends.rl_kernel_utils import ( + ExecutionDecision, + FallbackReason, + RlkOperatorUnavailable, + build_rlk_operator_adapter, + emit_execution_decision, + linear_logp_inputs_from_vime, +) +from vime.utils.rl_kernel import is_rl_kernel_op_enabled, is_rl_kernel_requested, is_rl_kernel_strict + +logger = logging.getLogger(__name__) + +_LINEAR_LOGP_ADAPTER = None +_LINEAR_LOGP_ADAPTER_ERROR: Exception | None = None +_WARNED_FALLBACK_REASONS: set[str] = set() +_FALLBACK_COUNTS: dict[str, int] = {"linear_logp": 0} +_LINEAR_LOGP_SAVE_PROBS_CAST_LOGGED = False +_RUNTIME_COUNTER_KEYS = ( + "linear_logp_call_count", + "linear_logp_token_count", + "linear_logp_dispatch_elapsed_s", + "linear_logp_fallback_count", +) +_RUNTIME_COUNTERS: dict[str, float] = dict.fromkeys(_RUNTIME_COUNTER_KEYS, 0.0) +_RUNTIME_COUNTER_LAST_SNAPSHOT: dict[str, float] = dict.fromkeys(_RUNTIME_COUNTER_KEYS, 0.0) +_NATIVE_LINEAR_LOGP_BACKEND = "vime.native.linear_logp" +_ZERO_TOKEN_LINEAR_LOGP_BACKEND = "vime.linear_logp.zero_tokens" +_RLK_LINEAR_LOGP_SUPPORT_PROVIDER = "rl_engine.kernels.support" +_RLK_LINEAR_LOGP_SUPPORT_UNAVAILABLE_ROW: dict[str, str] = { + "source": "vime_adapter", + "backend": "rl_kernel_unavailable", + "implementation": "RL-Kernel support provider is unavailable", + "dtype": "reported by RL-Kernel when installed", + "hardware": "reported by RL-Kernel when installed", + "tp": "reported by RL-Kernel when installed", + "cp": "reported by RL-Kernel when installed; vime decides CP fallback before dispatch", + "entropy": "linear_logp does not produce entropy; vime computes entropy on native fallback", + "full_gradient": "reported by RL-Kernel when installed", +} +_VIME_LINEAR_LOGP_SUPPORT_ROWS: tuple[dict[str, str], ...] = ( + { + "source": "vime_adapter", + "backend": "registry", + "implementation": "RlkRegistryOperatorAdapter.linear_logp -> kernel_registry.get_op('linear_logp')", + "dtype": "reported by RL-Kernel; vime records fp32 selected logprobs", + "hardware": "reported by the installed RL-Kernel backend", + "tp": "vime passes tp_group, vocab_start_index, and global_vocab_size when available", + "cp": "not supported; falls back before CP redistribution", + "entropy": "not supported; falls back when entropy is requested", + "full_gradient": "supported when the selected op returns an autograd-connected result", + }, + { + "source": "vime_adapter", + "backend": "native", + "implementation": "Megatron output layer + vime calculate_log_probs_and_entropy", + "dtype": "vime native logits path, fp32 logprob computation", + "hardware": "same as native vime/Megatron execution", + "tp": "supported by the native vime/Megatron logprob path", + "cp": "supported by the native vime/Megatron CP redistribution path", + "entropy": "supported by the native vime/Megatron path", + "full_gradient": "supported by native autograd over materialized logits", + }, +) + + +@dataclass +class LinearLogpRuntimeMetadata: + operator: str = "linear_logp" + requested_backend: str = "registry" + actual_backend: str = "not_selected" + backend_id: str | None = None + contract_id: str | None = None + fallback: bool = False + fallback_reason: str | None = None + memory_probe_enabled: bool = False + memory_alloc_delta_mb: float | None = None + memory_peak_alloc_delta_mb: float | None = None + memory_reserved_delta_mb: float | None = None + memory_peak_reserved_delta_mb: float | None = None + + +_LINEAR_LOGP_RUNTIME_METADATA = LinearLogpRuntimeMetadata() + + +@dataclass(frozen=True) +class LinearLogpContext: + lm_head_weight: torch.Tensor + bias: torch.Tensor | None + tp_group: Any + vocab_start_index: int = 0 + global_vocab_size: int | None = None + sequence_parallel: bool = False + + +def _env_flag(name: str) -> bool: + return os.getenv(name, "").strip().lower() in {"1", "true", "yes", "on"} + + +def _env_bool(name: str) -> bool | None: + value = os.getenv(name) + if value is None or value.strip() == "": + return None + lowered = value.strip().lower() + if lowered in {"1", "true", "yes", "on"}: + return True + if lowered in {"0", "false", "no", "off"}: + return False + raise ValueError(f"{name} must be a boolean flag, got {value!r}") + + +def _requested_linear_logp_backend() -> str: + requested = os.getenv("VIME_RL_KERNEL_LINEAR_LOGP_BACKEND", "").strip().lower() + aliases = { + "": "registry", + "auto": "registry", + "registry": "registry", + "rlk": "registry", + "rl_kernel": "registry", + } + return aliases.get(requested, requested) + + +def _stable_descriptor_id(value: str | None) -> float: + if not value: + return 0.0 + digest = hashlib.sha1(value.encode("utf-8")).hexdigest()[:12] + return float(int(digest, 16)) + + +def _fallback_code(reason: str) -> str: + code = re.sub(r"[^a-z0-9]+", "_", reason.lower()).strip("_") + return code[:80] or "linear_logp_fallback" + + +def _requested_modes(args: Namespace) -> tuple[str, str]: + config = getattr(args, "rlk_mode_config", None) + if config is not None: + return str(getattr(config, "fast", "off")), str(getattr(config, "consistency", "off")) + if getattr(args, "rlk_fast", None) is not None: + fast = str(args.rlk_fast) + elif getattr(args, "rl_kernel_strict", False): + fast = "strict" + elif getattr(args, "enable_rl_kernel", False): + fast = "auto" + else: + fast = "off" + return fast, str(getattr(args, "rlk_consistency", "off") or "off") + + +def _parallel_context() -> dict[str, Any]: + return { + "tp_world_size": int(mpu.get_tensor_model_parallel_world_size()), + "tp_rank": int(mpu.get_tensor_model_parallel_rank()), + "cp_world_size": int(mpu.get_context_parallel_world_size()), + "cp_rank": int(mpu.get_context_parallel_rank()), + } + + +def _emit_linear_logp_decision( + args: Namespace, + *, + decision: str, + actual_backend: str | None, + fallback: bool, + fallback_reason: str | None = None, + backend_id: str | None = None, + contract_id: str | None = None, + dtype: str | None = None, + details: dict[str, Any] | None = None, +) -> None: + fast, consistency = _requested_modes(args) + reason = None if fallback_reason is None else FallbackReason(code=_fallback_code(fallback_reason), message=fallback_reason) + record = ExecutionDecision( + operator="linear_logp", + stage="train_logprob", + requested_mode=f"fast={fast},consistency={consistency}", + requested_backend=_requested_linear_logp_backend(), + actual_backend=actual_backend, + decision=decision, + fallback=fallback, + fallback_reason=reason, + capability_backend_id=backend_id, + contract_id=contract_id, + dtype=dtype, + parallel_context=_parallel_context(), + strict_eligible=fast == "strict", + details=details or {}, + ) + emit_execution_decision(record) + + +def _reset_linear_logp_runtime_metadata() -> None: + global _LINEAR_LOGP_RUNTIME_METADATA + _LINEAR_LOGP_RUNTIME_METADATA = LinearLogpRuntimeMetadata( + requested_backend=_requested_linear_logp_backend(), + ) + + +def _clear_linear_logp_memory_metadata() -> None: + _LINEAR_LOGP_RUNTIME_METADATA.memory_probe_enabled = False + _LINEAR_LOGP_RUNTIME_METADATA.memory_alloc_delta_mb = None + _LINEAR_LOGP_RUNTIME_METADATA.memory_peak_alloc_delta_mb = None + _LINEAR_LOGP_RUNTIME_METADATA.memory_reserved_delta_mb = None + _LINEAR_LOGP_RUNTIME_METADATA.memory_peak_reserved_delta_mb = None + + +def _set_linear_logp_fallback(args: Namespace, reason: str) -> None: + _clear_linear_logp_memory_metadata() + _LINEAR_LOGP_RUNTIME_METADATA.requested_backend = _requested_linear_logp_backend() + _LINEAR_LOGP_RUNTIME_METADATA.actual_backend = _NATIVE_LINEAR_LOGP_BACKEND + _LINEAR_LOGP_RUNTIME_METADATA.backend_id = _NATIVE_LINEAR_LOGP_BACKEND + _LINEAR_LOGP_RUNTIME_METADATA.contract_id = "vime.native.linear_logp.selected_logprob" + _LINEAR_LOGP_RUNTIME_METADATA.fallback = True + _LINEAR_LOGP_RUNTIME_METADATA.fallback_reason = reason + _emit_linear_logp_decision( + args, + decision="fallback-native", + actual_backend=_NATIVE_LINEAR_LOGP_BACKEND, + fallback=True, + fallback_reason=reason, + backend_id=_NATIVE_LINEAR_LOGP_BACKEND, + contract_id=_LINEAR_LOGP_RUNTIME_METADATA.contract_id, + ) + + +def _set_linear_logp_selected_backend(args: Namespace, result: Any, dtype: torch.dtype) -> None: + _clear_linear_logp_memory_metadata() + decision = result.decision + provenance = dict(getattr(decision, "provenance", {}) or {}) + backend_id = provenance.get("backend_id") or getattr(decision, "backend", None) + contract_id = provenance.get("contract_id") + _LINEAR_LOGP_RUNTIME_METADATA.requested_backend = _requested_linear_logp_backend() + _LINEAR_LOGP_RUNTIME_METADATA.actual_backend = getattr(decision, "backend", None) + _LINEAR_LOGP_RUNTIME_METADATA.backend_id = backend_id + _LINEAR_LOGP_RUNTIME_METADATA.contract_id = contract_id + _LINEAR_LOGP_RUNTIME_METADATA.fallback = False + _LINEAR_LOGP_RUNTIME_METADATA.fallback_reason = None + _emit_linear_logp_decision( + args, + decision="optimized", + actual_backend=backend_id, + fallback=False, + backend_id=backend_id, + contract_id=contract_id, + dtype=str(dtype).replace("torch.", ""), + details={"implementation": getattr(decision, "backend", None)}, + ) + + +def _set_linear_logp_zero_token_decision(args: Namespace) -> None: + _clear_linear_logp_memory_metadata() + _LINEAR_LOGP_RUNTIME_METADATA.requested_backend = _requested_linear_logp_backend() + _LINEAR_LOGP_RUNTIME_METADATA.actual_backend = _ZERO_TOKEN_LINEAR_LOGP_BACKEND + _LINEAR_LOGP_RUNTIME_METADATA.backend_id = _ZERO_TOKEN_LINEAR_LOGP_BACKEND + _LINEAR_LOGP_RUNTIME_METADATA.contract_id = None + _LINEAR_LOGP_RUNTIME_METADATA.fallback = False + _LINEAR_LOGP_RUNTIME_METADATA.fallback_reason = None + _emit_linear_logp_decision( + args, + decision="optimized", + actual_backend=_ZERO_TOKEN_LINEAR_LOGP_BACKEND, + fallback=False, + backend_id=_ZERO_TOKEN_LINEAR_LOGP_BACKEND, + details={"zero_tokens": True}, + ) + + +def _record_linear_logp_memory_probe( + *, + alloc_before: int, + alloc_after: int, + peak_alloc: int, + reserved_before: int, + reserved_after: int, + peak_reserved: int, +) -> None: + mb = float(1024**2) + _LINEAR_LOGP_RUNTIME_METADATA.memory_probe_enabled = True + _LINEAR_LOGP_RUNTIME_METADATA.memory_alloc_delta_mb = (alloc_after - alloc_before) / mb + _LINEAR_LOGP_RUNTIME_METADATA.memory_peak_alloc_delta_mb = (peak_alloc - alloc_before) / mb + _LINEAR_LOGP_RUNTIME_METADATA.memory_reserved_delta_mb = (reserved_after - reserved_before) / mb + _LINEAR_LOGP_RUNTIME_METADATA.memory_peak_reserved_delta_mb = (peak_reserved - reserved_before) / mb + + +def get_linear_logp_support_matrix() -> tuple[dict[str, str], ...]: + return ( + *_load_rl_kernel_linear_logp_support_rows(), + *(dict(row) for row in _VIME_LINEAR_LOGP_SUPPORT_ROWS), + ) + + +def _load_rl_kernel_linear_logp_support_rows() -> tuple[dict[str, str], ...]: + try: + provider = importlib.import_module(_RLK_LINEAR_LOGP_SUPPORT_PROVIDER) + return tuple(dict(row) for row in provider.get_linear_logp_support_matrix()) + except Exception: + return (dict(_RLK_LINEAR_LOGP_SUPPORT_UNAVAILABLE_ROW),) + + +def get_linear_logp_runtime_metadata() -> dict[str, Any]: + metadata = asdict(_LINEAR_LOGP_RUNTIME_METADATA) + metadata["backend_descriptor_id"] = _stable_descriptor_id(metadata.get("backend_id")) + metadata["contract_descriptor_id"] = _stable_descriptor_id(metadata.get("contract_id")) + metadata["fallback_reason_descriptor_id"] = _stable_descriptor_id(metadata.get("fallback_reason")) + return metadata + + +def get_linear_logp_runtime_log_metrics(prefix: str = "train/rl_kernel_linear_logp_") -> dict[str, float]: + metadata = get_linear_logp_runtime_metadata() + metrics = { + f"{prefix}fallback": 1.0 if metadata.get("fallback") else 0.0, + f"{prefix}backend_descriptor_id": float(metadata["backend_descriptor_id"]), + f"{prefix}contract_descriptor_id": float(metadata["contract_descriptor_id"]), + f"{prefix}fallback_reason_descriptor_id": float(metadata["fallback_reason_descriptor_id"]), + } + for key in ( + "memory_alloc_delta_mb", + "memory_peak_alloc_delta_mb", + "memory_reserved_delta_mb", + "memory_peak_reserved_delta_mb", + ): + value = metadata.get(key) + if value is not None: + metrics[f"{prefix}{key}"] = float(value) + return metrics + + +def get_rl_kernel_fallback_count(op: str | None = None) -> int: + if op is not None: + return _FALLBACK_COUNTS.get(op, 0) + return sum(_FALLBACK_COUNTS.values()) + + +def reset_rl_kernel_runtime_counters() -> None: + for key in _RUNTIME_COUNTER_KEYS: + _RUNTIME_COUNTERS[key] = 0.0 + _RUNTIME_COUNTER_LAST_SNAPSHOT[key] = 0.0 + _reset_linear_logp_runtime_metadata() + + +def get_rl_kernel_runtime_counters() -> dict[str, float]: + return dict(_RUNTIME_COUNTERS) + + +def get_rl_kernel_runtime_counter_delta() -> dict[str, float]: + current = get_rl_kernel_runtime_counters() + delta = {key: current.get(key, 0.0) - _RUNTIME_COUNTER_LAST_SNAPSHOT.get(key, 0.0) for key in _RUNTIME_COUNTER_KEYS} + _RUNTIME_COUNTER_LAST_SNAPSHOT.update(current) + return delta + + +def _record_linear_logp_runtime(token_count: int, elapsed_s: float) -> None: + _RUNTIME_COUNTERS["linear_logp_call_count"] += 1.0 + _RUNTIME_COUNTERS["linear_logp_token_count"] += float(token_count) + _RUNTIME_COUNTERS["linear_logp_dispatch_elapsed_s"] += float(elapsed_s) + + +def _should_detach_linear_logp_hidden(args: Namespace) -> bool: + override = _env_bool("VIME_RL_KERNEL_LINEAR_LOGP_DETACH_HIDDEN") + if override is not None: + return override + patterns = tuple(getattr(args, "only_train_params_name_list", ()) or ()) + return bool(patterns) and all("output_layer" in str(pattern) for pattern in patterns) + + +def _linear_logp_needs_bf16_fast_path_cast() -> bool: + return _env_flag("RL_KERNEL_LINEAR_LOGP_SAVE_PROBS_BF16") or _env_flag("RL_KERNEL_LINEAR_LOGP_FUSED_TILE_BWD_FULL") + + +def _maybe_cast_hidden_for_bf16_fast_path(hidden_states: torch.Tensor, weight: torch.Tensor) -> torch.Tensor: + global _LINEAR_LOGP_SAVE_PROBS_CAST_LOGGED + if not _linear_logp_needs_bf16_fast_path_cast(): + return hidden_states + if not (hidden_states.is_cuda and weight.is_cuda and hidden_states.device == weight.device): + return hidden_states + if weight.dtype != torch.bfloat16 or hidden_states.dtype == torch.bfloat16: + return hidden_states + if not hidden_states.is_floating_point(): + return hidden_states + + if not _LINEAR_LOGP_SAVE_PROBS_CAST_LOGGED: + logger.info( + "Casting RL-Kernel linear_logp hidden states from %s to bf16 to enable bf16 fast path.", + hidden_states.dtype, + ) + _LINEAR_LOGP_SAVE_PROBS_CAST_LOGGED = True + return hidden_states.to(dtype=torch.bfloat16) + + +def _warn_fallback(args: Namespace, reason: str) -> None: + _FALLBACK_COUNTS["linear_logp"] = _FALLBACK_COUNTS.get("linear_logp", 0) + 1 + _RUNTIME_COUNTERS["linear_logp_fallback_count"] += 1.0 + _set_linear_logp_fallback(args, reason) + if is_rl_kernel_strict(args): + raise RuntimeError(f"RL-Kernel linear_logp is enabled but unavailable: {reason}") + if reason not in _WARNED_FALLBACK_REASONS: + logger.warning("Falling back to vime logprob path because RL-Kernel linear_logp is unavailable: %s", reason) + _WARNED_FALLBACK_REASONS.add(reason) + + +def _get_linear_logp_adapter(args: Namespace): + global _LINEAR_LOGP_ADAPTER, _LINEAR_LOGP_ADAPTER_ERROR + if _LINEAR_LOGP_ADAPTER is not None: + return _LINEAR_LOGP_ADAPTER + if _LINEAR_LOGP_ADAPTER_ERROR is not None: + _warn_fallback(args, str(_LINEAR_LOGP_ADAPTER_ERROR)) + return None + + try: + _LINEAR_LOGP_ADAPTER = build_rlk_operator_adapter(args=args, backend="auto") + logger.info("Using RL-Kernel operator adapter for linear_logp: %s", type(_LINEAR_LOGP_ADAPTER).__name__) + return _LINEAR_LOGP_ADAPTER + except Exception as exc: # pragma: no cover - exercised when optional dependencies are unavailable + _LINEAR_LOGP_ADAPTER_ERROR = exc + _warn_fallback(args, str(exc)) + return None + + +def _unwrap_model_chunk(model): + while hasattr(model, "module"): + model = model.module + return model + + +def _is_pipeline_last_stage_for_model(model) -> bool: + module = _unwrap_model_chunk(model) + vp_stage = getattr(module, "vp_stage", None) + try: + vp_world_size = mpu.get_virtual_pipeline_model_parallel_world_size() + except Exception: + vp_world_size = None + + try: + if vp_world_size is not None and vp_stage is not None: + return bool(mpu.is_pipeline_last_stage(ignore_virtual=False, vp_stage=vp_stage)) + return bool(mpu.is_pipeline_last_stage(ignore_virtual=True)) + except Exception: + return True + + +def _get_lm_head_weight(model, output_layer) -> torch.Tensor | None: + if getattr(model, "share_embeddings_and_output_weights", False): + shared_weight = getattr(model, "shared_embedding_or_output_weight", None) + if callable(shared_weight): + try: + weight = shared_weight() + if isinstance(weight, torch.Tensor): + return weight + except Exception: + logger.debug("Unable to read shared embedding/output weight for RL-Kernel linear_logp.", exc_info=True) + + weight = getattr(output_layer, "weight", None) + if isinstance(weight, torch.Tensor): + return weight + + shared_weight = getattr(model, "shared_embedding_or_output_weight", None) + if callable(shared_weight): + try: + weight = shared_weight() + if isinstance(weight, torch.Tensor): + return weight + except Exception: + logger.debug("Unable to read shared embedding/output weight for RL-Kernel linear_logp.", exc_info=True) + + return None + + +def get_linear_logp_context_from_model(args: Namespace, model) -> LinearLogpContext | None: + if not is_rl_kernel_op_enabled(args, "linear_logp"): + return None + + if not _is_pipeline_last_stage_for_model(model): + return None + + module = _unwrap_model_chunk(model) + output_layer = getattr(module, "output_layer", None) + if output_layer is None: + _warn_fallback(args, "model output_layer is unavailable") + return None + + weight = _get_lm_head_weight(module, output_layer) + if weight is None: + _warn_fallback(args, "LM-head weight is unavailable") + return None + + bias = getattr(output_layer, "bias", None) + if not isinstance(bias, torch.Tensor): + bias = None + + tp_world_size = int(mpu.get_tensor_model_parallel_world_size()) + tp_group = mpu.get_tensor_model_parallel_group() if tp_world_size > 1 else None + vocab_start_index = 0 + global_vocab_size = None + if tp_world_size > 1: + local_vocab_size = int(weight.size(0)) + vocab_start_index = int(mpu.get_tensor_model_parallel_rank()) * local_vocab_size + global_vocab_size = getattr(args, "padded_vocab_size", None) + if global_vocab_size is None: + global_vocab_size = local_vocab_size * tp_world_size + + return LinearLogpContext( + lm_head_weight=weight, + bias=bias, + tp_group=tp_group, + vocab_start_index=vocab_start_index, + global_vocab_size=None if global_vocab_size is None else int(global_vocab_size), + sequence_parallel=bool(getattr(output_layer, "sequence_parallel", getattr(args, "sequence_parallel", False))), + ) + + +def _linear_logp_runtime_blocker(args: Namespace, *, with_entropy: bool) -> str | None: + if with_entropy: + return "entropy is requested" + if getattr(args, "qkv_format", "thd") != "thd": + return "only qkv_format=thd is supported by RL-Kernel linear_logp" + if mpu.get_context_parallel_world_size() != 1 or getattr(args, "allgather_cp", False): + return "context parallel logprob redistribution is not supported by RL-Kernel linear_logp" + if getattr(args, "rollout_temperature", 1.0) <= 0: + return "rollout_temperature must be positive" + return None + + +def should_use_linear_logp_model_output(args: Namespace, *, with_entropy: bool) -> bool: + if not is_rl_kernel_op_enabled(args, "linear_logp"): + return False + reason = _linear_logp_runtime_blocker(args, with_entropy=with_entropy) + if reason is not None: + _warn_fallback(args, reason) + return False + return True + + +def warn_linear_logp_fallback(args: Namespace, reason: str) -> None: + if is_rl_kernel_requested(args): + _warn_fallback(args, reason) + + +@contextmanager +def return_hidden_states_for_linear_logp(args: Namespace, model, context: LinearLogpContext | None): + if context is None: + yield False + return + + module = _unwrap_model_chunk(model) + if not hasattr(module, "post_process"): + _warn_fallback(args, "model post_process flag is unavailable") + yield False + return + + old_post_process = module.post_process + module.post_process = False + try: + yield True + finally: + module.post_process = old_post_process + + +def maybe_compute_linear_logp( + hidden_states: torch.Tensor, + target_ids: torch.Tensor, + *, + context: LinearLogpContext | None, + args: Namespace, + with_entropy: bool, +) -> torch.Tensor | None: + if not is_rl_kernel_op_enabled(args, "linear_logp"): + return None + + reason = _linear_logp_runtime_blocker(args, with_entropy=with_entropy) + if reason is not None: + _warn_fallback(args, reason) + return None + + if context is None: + _warn_fallback(args, "hidden-state linear_logp context is unavailable") + return None + + if target_ids.numel() == 0: + _set_linear_logp_zero_token_decision(args) + return hidden_states.new_zeros((0,), dtype=torch.float32) + + adapter = _get_linear_logp_adapter(args) + if adapter is None: + return None + + weight = context.lm_head_weight + bias = context.bias + rollout_temperature = float(getattr(args, "rollout_temperature", 1.0)) + if rollout_temperature != 1.0: + weight = weight / rollout_temperature + if bias is not None: + bias = bias / rollout_temperature + if _should_detach_linear_logp_hidden(args): + hidden_states = hidden_states.detach() + hidden_states = _maybe_cast_hidden_for_bf16_fast_path(hidden_states, weight) + + memory_probe = _env_flag("VIME_LINEAR_LOGP_MEMORY_PROBE") and hidden_states.is_cuda + if memory_probe: + probe_device = hidden_states.device + torch.cuda.synchronize(probe_device) + probe_before_alloc = torch.cuda.memory_allocated(probe_device) + probe_before_reserved = torch.cuda.memory_reserved(probe_device) + torch.cuda.reset_peak_memory_stats(probe_device) + + start_s = time.perf_counter() + try: + result = adapter.linear_logp( + linear_logp_inputs_from_vime( + hidden=hidden_states, + lm_head_weight=weight, + target_ids=target_ids.long(), + bias=bias, + tp_group=context.tp_group, + vocab_start_index=context.vocab_start_index, + global_vocab_size=context.global_vocab_size, + metadata={"requested_backend": _requested_linear_logp_backend()}, + ) + ) + except RlkOperatorUnavailable: + raise + except Exception as exc: + _warn_fallback(args, str(exc)) + return None + + elapsed_s = getattr(result.decision, "elapsed_s", 0.0) or (time.perf_counter() - start_s) + if result.value is None: + _warn_fallback(args, result.decision.reason or "adapter returned no linear_logp value") + return None + + _set_linear_logp_selected_backend(args, result, hidden_states.dtype) + _record_linear_logp_runtime(target_ids.numel(), elapsed_s) + + if memory_probe: + torch.cuda.synchronize(probe_device) + probe_after_alloc = torch.cuda.memory_allocated(probe_device) + probe_after_reserved = torch.cuda.memory_reserved(probe_device) + probe_peak_alloc = torch.cuda.max_memory_allocated(probe_device) + probe_peak_reserved = torch.cuda.max_memory_reserved(probe_device) + _record_linear_logp_memory_probe( + alloc_before=probe_before_alloc, + alloc_after=probe_after_alloc, + peak_alloc=probe_peak_alloc, + reserved_before=probe_before_reserved, + reserved_after=probe_after_reserved, + peak_reserved=probe_peak_reserved, + ) + logger.info( + "RL-Kernel linear_logp memory_probe: hidden_shape=%s weight_shape=%s tokens=%d alloc_delta_mb=%.2f peak_alloc_delta_mb=%.2f reserved_delta_mb=%.2f peak_reserved_delta_mb=%.2f", + tuple(hidden_states.shape), + tuple(weight.shape), + int(target_ids.numel()), + _LINEAR_LOGP_RUNTIME_METADATA.memory_alloc_delta_mb, + _LINEAR_LOGP_RUNTIME_METADATA.memory_peak_alloc_delta_mb, + _LINEAR_LOGP_RUNTIME_METADATA.memory_reserved_delta_mb, + _LINEAR_LOGP_RUNTIME_METADATA.memory_peak_reserved_delta_mb, + ) + + return result.value.float().reshape(-1) diff --git a/vime/backends/rl_kernel_utils/adapter.py b/vime/backends/rl_kernel_utils/adapter.py index 333cede4..6af15481 100644 --- a/vime/backends/rl_kernel_utils/adapter.py +++ b/vime/backends/rl_kernel_utils/adapter.py @@ -471,15 +471,31 @@ def _call(self, op_name: str, token_count: int | None, fn: Callable[[Any], Any]) 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) + try: + value = fn(op) + except Exception as exc: + capability = RlkCapability( + op_name=op_name, + available=False, + backend=type(op).__name__, + reason=f"RL-Kernel op {op_name!r} failed during execution: {exc}", + ) + return self._unsupported_result_or_raise(op_name, capability, token_count=token_count) elapsed_s = time.perf_counter() - start + backend_id = _op_text_attr(op, "backend_id", "backend_name", "name") + contract_id = _op_text_attr(op, "contract_id", "numeric_contract_id") + provenance = dict(self.provenance()) + if backend_id is not None: + provenance["backend_id"] = backend_id + if contract_id is not None: + provenance["contract_id"] = contract_id decision = RlkOperatorDecision( op_name=op_name, path="fast", backend=type(op).__name__, elapsed_s=elapsed_s, token_count=token_count, - provenance=self.provenance(), + provenance=provenance, ) self.telemetry.record_decision(decision) return RlkOperatorResult(value=value, decision=decision) @@ -678,6 +694,21 @@ def _registry_op_name(op_name: str) -> str: return RLK_OP_SELECTED_LOGPROBS if op_name == RLK_OP_REFERENCE_LOGPROBS else op_name +def _op_text_attr(op: Any, *names: str) -> str | None: + for name in names: + value = getattr(op, name, None) + if value is None: + continue + if callable(value): + try: + value = value() + except TypeError: + continue + if value is not None: + return str(value) + return None + + def _mask_inactive_values(value: Any, mask: Any | None) -> Any: if mask is None: return value diff --git a/vime/utils/rl_kernel.py b/vime/utils/rl_kernel.py new file mode 100644 index 00000000..a666050b --- /dev/null +++ b/vime/utils/rl_kernel.py @@ -0,0 +1,51 @@ +from __future__ import annotations + +from argparse import Namespace +from collections.abc import Iterable + + +def _normalize_ops(value: str | Iterable[str] | None) -> tuple[str, ...]: + if value is None: + return () + if isinstance(value, str): + raw_items = value.replace(",", " ").split() + else: + raw_items = [] + for item in value: + raw_items.extend(str(item).replace(",", " ").split()) + return tuple(dict.fromkeys(op.strip() for op in raw_items if op.strip())) + + +def _resolved_fast_mode(args: Namespace) -> str: + config = getattr(args, "rlk_mode_config", None) + if config is not None: + return str(getattr(config, "fast", "off")).lower() + if getattr(args, "rlk_fast", None) is not None: + return str(args.rlk_fast).lower() + if getattr(args, "rl_kernel_strict", False): + return "strict" + if getattr(args, "enable_rl_kernel", False): + return "auto" + return "off" + + +def _resolved_ops(args: Namespace) -> tuple[str, ...]: + config = getattr(args, "rlk_mode_config", None) + if config is not None: + return _normalize_ops(getattr(config, "ops", ())) + return _normalize_ops(getattr(args, "rl_kernel_ops", ())) + + +def is_rl_kernel_requested(args: Namespace) -> bool: + return _resolved_fast_mode(args) != "off" + + +def is_rl_kernel_strict(args: Namespace) -> bool: + return _resolved_fast_mode(args) == "strict" + + +def is_rl_kernel_op_enabled(args: Namespace, op: str) -> bool: + if not is_rl_kernel_requested(args): + return False + ops = _resolved_ops(args) + return "*" in ops or op in ops