diff --git a/tests/_unit_stubs.py b/tests/_unit_stubs.py index 475db552..7f64320c 100644 --- a/tests/_unit_stubs.py +++ b/tests/_unit_stubs.py @@ -47,15 +47,31 @@ def ensure_ray_stub() -> None: if real_module_available("ray"): return ray = MagicMock() + ray.__path__ = [] + ray.get = lambda refs: refs + ray.put = lambda value, **kwargs: value # noqa: ARG005 + ray.remote = _ray_remote_stub + ray.util = types.ModuleType("ray.util") + ray.util.__path__ = [] + scheduling_strategies = types.ModuleType("ray.util.scheduling_strategies") + scheduling_strategies.PlacementGroupSchedulingStrategy = type( + "PlacementGroupSchedulingStrategy", + (), + {"__init__": lambda self, **kwargs: setattr(self, "kwargs", kwargs)}, + ) + ray.util.scheduling_strategies = scheduling_strategies sys.modules["ray"] = ray sys.modules["ray._private"] = MagicMock() sys.modules["ray._private.services"] = MagicMock() sys.modules["ray.actor"] = MagicMock() + sys.modules["ray.util"] = ray.util + sys.modules["ray.util.scheduling_strategies"] = scheduling_strategies def install_rollout_optional_stubs() -> None: """Stub rollout-side optional imports when not installed.""" ensure_ray_stub() + install_pyarrow_stub() install_vllm_router_stub() @@ -157,6 +173,24 @@ def install_wandb_stub() -> None: sys.modules["wandb"] = wandb_mod +def install_pyarrow_stub() -> None: + if "pyarrow" in sys.modules: + return + + pyarrow_mod = types.ModuleType("pyarrow") + pyarrow_mod.__path__ = [] + parquet_mod = types.ModuleType("pyarrow.parquet") + + class ParquetFile: + def __init__(self, *args, **kwargs): # noqa: ARG002 + raise ImportError("pyarrow parquet support is unavailable in this unit-test environment") + + parquet_mod.ParquetFile = ParquetFile + pyarrow_mod.parquet = parquet_mod + sys.modules["pyarrow"] = pyarrow_mod + sys.modules["pyarrow.parquet"] = parquet_mod + + def save_sys_modules(names: Iterable[str]) -> dict[str, Any]: return {k: sys.modules.get(k) for k in names} @@ -221,13 +255,48 @@ def install_megatron_mpu_stub() -> MagicMock: def install_ray_stub() -> None: ray_mod = types.ModuleType("ray") + ray_mod.__path__ = [] ray_mod.get = lambda refs: refs + ray_mod.put = lambda value, **kwargs: value # noqa: ARG005 + ray_mod.remote = _ray_remote_stub ray_mod.ObjectRef = object ray_mod.actor = types.ModuleType("ray.actor") ray_mod.actor.ActorHandle = object ray_mod._private = types.SimpleNamespace(services=types.SimpleNamespace(get_node_ip_address=lambda: "127.0.0.1")) + ray_mod.util = types.ModuleType("ray.util") + ray_mod.util.__path__ = [] + scheduling_strategies = types.ModuleType("ray.util.scheduling_strategies") + scheduling_strategies.PlacementGroupSchedulingStrategy = type( + "PlacementGroupSchedulingStrategy", + (), + {"__init__": lambda self, **kwargs: setattr(self, "kwargs", kwargs)}, + ) + ray_mod.util.scheduling_strategies = scheduling_strategies sys.modules.setdefault("ray", ray_mod) sys.modules.setdefault("ray.actor", ray_mod.actor) + sys.modules.setdefault("ray.util", ray_mod.util) + sys.modules.setdefault("ray.util.scheduling_strategies", scheduling_strategies) + + +class _RayRemoteWrapper: + def __init__(self, target): + self.__ray_actor_class__ = target + self._target = target + + def options(self, **kwargs): # noqa: ARG002 + return self + + def remote(self, *args, **kwargs): + return self._target(*args, **kwargs) + + +def _ray_remote_stub(target=None, **kwargs): # noqa: ARG001 + def wrap(obj): + return _RayRemoteWrapper(obj) if isinstance(obj, type) else obj + + if target is None: + return wrap + return wrap(target) def install_vllm_cli_stubs() -> None: diff --git a/tests/utils/test_consistency_audit.py b/tests/utils/test_consistency_audit.py new file mode 100644 index 00000000..8f3edaba --- /dev/null +++ b/tests/utils/test_consistency_audit.py @@ -0,0 +1,326 @@ +from __future__ import annotations + +from argparse import Namespace + +import pytest +import torch + +from vime.utils.consistency_audit import ( + build_batch_invariance_replay_cases, + build_consistency_diagnostic_metadata, + build_consistency_replay_manifest, + run_consistency_audit, + validate_consistency_audit_batch, +) +from vime.utils.consistency_metadata import build_rollout_consistency_metadata, stable_fingerprint +from vime.utils.types import Sample + +NUM_GPUS = 0 + + +def _args(mode: str = "audit", **overrides) -> Namespace: + values = dict( + rlk_consistency=mode, + model_name="unit/model", + train_backend="megatron", + hf_checkpoint="unit/model", + padding_side="right", + quantization="none", + tensor_model_parallel_size=1, + context_parallel_size=1, + sequence_parallel=False, + router_policy="consistent_hash", + vllm_enable_prefix_caching=False, + vllm_enable_deterministic_inference=True, + params_dtype="bf16", + eps_clip=0.2, + ) + values.update(overrides) + return Namespace(**values) + + +def _sample(index: int = 5, **overrides) -> Sample: + values = dict( + index=index, + group_index=2, + rollout_id=index + 10, + session_id=f"session-{index}", + tokens=[101, 201, 202], + response_length=2, + loss_mask=[1, 1], + rollout_log_probs=[-0.1, -0.2], + weight_versions=["weights-v1"], + status=Sample.Status.COMPLETED, + metadata={ + "pre_update": True, + "position_cache": {"position_ids": [0, 1, 2], "cache_policy": "none"}, + "quantization": {"policy": "none"}, + }, + ) + values.update(overrides) + return Sample(**values) + + +def _metadata(sample: Sample | None = None, *, args: Namespace | None = None, **overrides): + sample = sample or _sample() + record = build_rollout_consistency_metadata( + sample, + args=args or _args(), + sampling_params={"temperature": 1.0, "top_p": 1.0, "top_k": 0}, + logprob_contract_id="rlk.logp.native.fp32", + requested_provenance={"backend": "native", "fallback": False}, + actual_provenance={"backend": "native", "fallback": False}, + ) + record.update(overrides) + return record + + +def _batch(*, metadata=None, layout=None): + metadata = [_metadata()] if metadata is None else metadata + layout = [ + { + "fingerprint": "layout-1", + "active_mask_density": 1.0, + "dp_rank": 0, + "microbatch_id": 0, + "microbatch_offset": 0, + } + ] if layout is None else layout + return { + "tokens": [torch.tensor([101, 201, 202])], + "unconcat_tokens": [torch.tensor([101, 201, 202])], + "total_lengths": [3], + "response_lengths": [2], + "loss_masks": [torch.tensor([1, 1])], + "rollout_log_probs": [torch.tensor([-0.1, -0.2])], + "sample_indices": [5], + "rollout_ids": [15], + "consistency_metadata": metadata, + "consistency_batch_layout_fingerprints": layout, + } + + +@pytest.mark.unit +def test_diagnostic_metadata_prefers_consistency_records_and_batch_layouts(): + metadata = build_consistency_diagnostic_metadata(_batch()) + + assert metadata[0]["model_name"] == "unit/model" + assert metadata[0]["backend_id"] == "native" + assert metadata[0]["contract_id"] == "rlk.logp.native.fp32" + assert metadata[0]["batch_layout_fingerprint"] == "layout-1" + assert metadata[0]["provenance_fingerprint"].startswith("sha256:") + + +@pytest.mark.unit +def test_run_consistency_audit_builds_metrics_manifest_and_result_cube(): + batch = _batch() + + result = run_consistency_audit( + [torch.tensor([-0.1, -0.1])], + batch["rollout_log_probs"], + batch["loss_masks"], + args=_args("audit"), + batch=batch, + rank=3, + ) + + assert result.metadata_validation.ok + assert result.metrics["rlk_audit_active_token_count"].item() == pytest.approx(2.0) + assert result.metrics["rlk_audit_dlogp_abs_max"].item() == pytest.approx(0.1) + assert result.metrics["rlk_audit_metadata_warning_count"].item() == pytest.approx(0.0) + assert result.metrics["rlk_audit_replay_case_count"].item() == pytest.approx(5.0) + assert result.replay_manifest["rank"] == 3 + assert result.replay_manifest["samples"][0]["batch_layout_fingerprint"] == "layout-1" + assert result.result_cube["axes"]["batch_layout"] == "layout-1" + assert result.result_cube["axes"]["dtype"] == "bf16" + assert result.result_cube["axes"]["tp"] == 1 + assert result.result_cube["axes"]["logp_backend"] == "native" + assert result.result_cube["metrics"]["max_abs_dlogp"] == pytest.approx(0.1) + + +@pytest.mark.unit +def test_run_consistency_audit_prefers_runtime_provenance_for_result_cube(): + runtime_provenance = { + "operator": "linear_logp", + "requested_backend": "registry", + "actual_backend": "rl_engine.linear_logp", + "backend_id": "rlk.linear_logp.fast", + "contract_id": "rlk.linear_logp.fp32", + "fallback": False, + "strict_failure": False, + } + + result = run_consistency_audit( + [torch.tensor([-0.1, -0.2])], + [torch.tensor([-0.1, -0.2])], + [torch.tensor([1, 1])], + args=_args("audit", params_dtype=torch.bfloat16), + batch=_batch(), + runtime_provenance=runtime_provenance, + ) + + assert result.metrics["rlk_audit_runtime_fallback"].item() == pytest.approx(0.0) + assert result.replay_manifest["runtime_provenance"] == runtime_provenance + assert result.replay_manifest["runtime_provenance_fingerprint"].startswith("sha256:") + assert result.result_cube["axes"]["dtype"] == "bfloat16" + assert result.result_cube["axes"]["logp_backend"] == "rlk.linear_logp.fast" + assert result.result_cube["runtime_provenance"] == runtime_provenance + + +@pytest.mark.unit +def test_strict_audit_rejects_missing_position_cache_before_drift_attribution(): + record = _metadata() + record["position_cache"] = {"fingerprint": None} + batch = _batch(metadata=[record]) + + with pytest.raises(ValueError, match="position_cache_metadata_missing"): + run_consistency_audit( + [torch.tensor([-0.1, -0.2])], + batch["rollout_log_probs"], + batch["loss_masks"], + args=_args("strict"), + batch=batch, + ) + + +@pytest.mark.unit +def test_strict_audit_requires_batch_layout_before_drift_attribution(): + batch = _batch(layout=[]) + + with pytest.raises(ValueError, match="batch_layout_missing"): + run_consistency_audit( + [torch.tensor([-0.1, -0.2])], + batch["rollout_log_probs"], + batch["loss_masks"], + args=_args("strict"), + batch=batch, + ) + + +@pytest.mark.unit +def test_audit_mode_reports_missing_quantization_as_metadata_warning(): + record = _metadata() + record["quantization"] = {"fingerprint": None} + batch = _batch(metadata=[record]) + + validation = validate_consistency_audit_batch(batch, mode="audit") + + assert validation.ok + assert [issue.code for issue in validation.warnings] == ["quantization_metadata_missing"] + + +@pytest.mark.unit +def test_runtime_fallback_is_warning_in_audit_and_failure_in_strict(): + runtime_provenance = { + "operator": "linear_logp", + "requested_backend": "registry", + "actual_backend": "vime.native.linear_logp", + "fallback": True, + "fallback_reason": "unit fallback", + } + + validation = validate_consistency_audit_batch( + _batch(), + mode="audit", + runtime_provenance=runtime_provenance, + ) + + assert validation.ok + assert [issue.code for issue in validation.warnings] == ["undeclared_linear_logp_runtime_fallback"] + + with pytest.raises(ValueError, match="undeclared_linear_logp_runtime_fallback"): + run_consistency_audit( + [torch.tensor([-0.1, -0.2])], + [torch.tensor([-0.1, -0.2])], + [torch.tensor([1, 1])], + args=_args("strict"), + batch=_batch(), + runtime_provenance=runtime_provenance, + ) + + +@pytest.mark.unit +def test_audit_batch_missing_metadata_warning_is_separate_from_dlogp_warning(): + batch = { + "rollout_log_probs": [torch.tensor([0.0])], + "loss_masks": [torch.tensor([1])], + "sample_indices": [1], + "rollout_ids": [2], + } + + result = run_consistency_audit( + [torch.tensor([0.25])], + batch["rollout_log_probs"], + batch["loss_masks"], + args=_args("audit", rlk_contract_id="contract", rlk_batch_layout_fingerprint="layout"), + batch=batch, + model_name="unit/model", + backend_id="megatron", + provenance_fingerprint="prov", + ) + + assert result.metrics["rlk_audit_warning_count"].item() == pytest.approx(0.0) + assert result.metrics["rlk_audit_metadata_warning_count"].item() == pytest.approx(1.0) + assert result.metadata_validation.warnings[0].code == "consistency_metadata_missing" + + +@pytest.mark.unit +def test_replay_manifest_treats_1d_tensor_fields_as_one_sample(): + batch = { + "response_lengths": torch.tensor([2]), + "total_lengths": torch.tensor([3]), + "loss_masks": torch.tensor([1, 0]), + "rollout_log_probs": torch.tensor([-0.1, -0.2]), + } + + manifest = build_consistency_replay_manifest(batch, mode="audit", rank=0) + + assert manifest["sample_count"] == 1 + assert manifest["samples"][0]["response_length"] == 2 + assert manifest["samples"][0]["active_token_count"] == 1 + assert manifest["samples"][0]["has_rollout_log_probs"] is True + assert len(manifest["batch_invariance_cases"]) == 5 + + +@pytest.mark.unit +def test_replay_manifest_contains_batch_invariance_cases_without_tensor_payloads(): + batch = _batch() + batch["consistency_metadata"][0]["dynamic_sampling"] = {"keep": True, "reason": "unit"} + batch["consistency_metadata"][0]["fingerprint"] = stable_fingerprint(batch["consistency_metadata"][0]) + + manifest = build_consistency_replay_manifest(batch, mode="audit", rank=0) + cases = build_batch_invariance_replay_cases(batch) + + assert manifest["sample_count"] == 1 + assert manifest["samples"][0]["has_rollout_log_probs"] is True + assert manifest["samples"][0]["dynamic_sampling"] == {"keep": True, "reason": "unit"} + assert {case["case"] for case in cases} == { + "same_sample_alone", + "same_sample_mixed_batch", + "padding_packing_variant", + "active_token_density_variant", + "dynamic_sampling_keep_drop_variant", + } + assert all("tokens" not in sample for sample in manifest["samples"]) + assert manifest["fingerprint"].startswith("sha256:") + + +@pytest.mark.unit +def test_debug_train_dump_preserves_consistency_replay_manifest(tmp_path, monkeypatch): + from vime.utils import train_dump_utils + + manifest = build_consistency_replay_manifest(_batch(), mode="audit", rank=0) + path_template = str(tmp_path / "train_{rollout_id}_{rank}.pt") + args = Namespace(save_debug_train_data=path_template) + monkeypatch.setattr(torch.distributed, "get_rank", lambda: 0) + + train_dump_utils.save_debug_train_data( + args, + rollout_id=12, + rollout_data={"consistency_replay_manifest": manifest}, + ) + + payload = torch.load(tmp_path / "train_12_0.pt", weights_only=False) + assert payload["rollout_id"] == 12 + assert payload["rank"] == 0 + assert payload["rollout_data"]["consistency_replay_manifest"]["fingerprint"] == manifest["fingerprint"] diff --git a/tests/utils/test_consistency_metadata.py b/tests/utils/test_consistency_metadata.py index fd1c6c81..c592bd82 100644 --- a/tests/utils/test_consistency_metadata.py +++ b/tests/utils/test_consistency_metadata.py @@ -1,11 +1,21 @@ from __future__ import annotations import argparse +import sys +from pathlib import Path import pytest -from vime.rollout.data_source import RolloutDataSourceWithBuffer -from vime.utils.consistency_metadata import ( +_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 # noqa: E402 + +_unit_stubs.install_rollout_optional_stubs() + +from vime.rollout.data_source import RolloutDataSourceWithBuffer # noqa: E402 +from vime.utils.consistency_metadata import ( # noqa: E402 build_batch_layout_fingerprints, build_requested_actual_provenance, build_rollout_consistency_metadata, @@ -15,7 +25,7 @@ stable_fingerprint, validate_samples_consistency_metadata, ) -from vime.utils.types import Sample +from vime.utils.types import Sample # noqa: E402 NUM_GPUS = 0 diff --git a/tests/utils/test_consistency_metadata_rollout_integration.py b/tests/utils/test_consistency_metadata_rollout_integration.py index 146611d8..b3303988 100644 --- a/tests/utils/test_consistency_metadata_rollout_integration.py +++ b/tests/utils/test_consistency_metadata_rollout_integration.py @@ -14,11 +14,11 @@ _unit_stubs.install_rollout_optional_stubs() _unit_stubs.install_vllm_cli_stubs() +import vime.ray.rollout as rollout_mod # noqa: E402 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 @@ -35,6 +35,13 @@ class Args: rollout_num_gpus = 1 rollout_num_gpus_per_engine = 1 router_policy = "round_robin" + global_batch_size = 1 + micro_batch_size = 1 + use_dynamic_batch_size = False + max_tokens_per_gpu = None + balance_data = False + balance_by_flops = False + rollout_data_transport = "object-store" def _manager(mode: str): @@ -104,3 +111,45 @@ def test_convert_samples_to_train_data_carries_sample_metadata_when_present(): 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 + + +@pytest.mark.unit +def test_split_train_data_attaches_consistency_replay_manifest(monkeypatch): + sample = _sample() + sample.metadata["position_cache"] = {"position_ids": [0, 1, 2], "cache_policy": "none"} + sample.metadata["quantization"] = {"policy": "none"} + 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}, + ) + manager = _manager("audit") + manager.train_parallel_config = { + "dp_size": 1, + "cp_size": 1, + "vpp_size": 1, + "microbatch_group_size_per_vp_stage": 1, + } + monkeypatch.setattr(rollout_mod.ray, "put", lambda value, **kwargs: value) + + train_data = manager._convert_samples_to_train_data([sample]) + refs = manager._split_train_data_by_dp(train_data) + rollout_data = refs[0].inner + manifest = rollout_data["consistency_replay_manifest"] + + assert manifest["mode"] == "audit" + assert manifest["rank"] == 0 + assert manifest["sample_count"] == 1 + assert manifest["samples"][0]["sample_index"] == 0 + assert manifest["samples"][0]["rollout_id"] == 0 + assert manifest["samples"][0]["has_rollout_log_probs"] is True + assert {case["case"] for case in manifest["batch_invariance_cases"]} == { + "same_sample_alone", + "same_sample_mixed_batch", + "padding_packing_variant", + "active_token_density_variant", + "dynamic_sampling_keep_drop_variant", + } diff --git a/tests/utils/test_dlogp_diagnostics.py b/tests/utils/test_dlogp_diagnostics.py index 507ea8ed..ca5d09b5 100644 --- a/tests/utils/test_dlogp_diagnostics.py +++ b/tests/utils/test_dlogp_diagnostics.py @@ -15,6 +15,7 @@ import _unit_stubs +from vime.utils.consistency_metadata import stable_fingerprint from vime.utils.dlogp_diagnostics import compute_dlogp_diagnostics, get_rlk_consistency_mode, is_dlogp_audit_enabled NUM_GPUS = 0 @@ -98,6 +99,31 @@ def _policy_batch() -> dict: } +def _complete_consistency_record() -> dict: + record = { + "schema_version": 1, + "sample": {"index": 42, "rollout_id": 7, "session_id": "session-42"}, + "tokens": {"response_token_ids_fingerprint": stable_fingerprint([1, 2])}, + "active_mask": {"mask_fingerprint": stable_fingerprint([1, 1]), "active_token_count": 2}, + "tokenizer": {"fingerprint": "tokenizer"}, + "sampling": {"params_fingerprint": "sampling"}, + "padding": {"side": "right"}, + "position_cache": {"fingerprint": "position-cache"}, + "quantization": {"fingerprint": "quantization"}, + "model": {"name": "record-model"}, + "weight": {"version": "weights-v1", "pre_update": True}, + "old_logp": {"source": "rollout_engine", "contract_id": "record-contract"}, + "provenance": { + "actual": {"backend": "record-backend", "fallback": False}, + "actual_fingerprint": "record-provenance", + "mismatches": {}, + "undeclared_fallback": False, + }, + } + record["fingerprint"] = stable_fingerprint(record) + return record + + @pytest.mark.unit def test_dlogp_metrics_use_active_tokens_only_and_identify_worst_token(): train_log_probs = [ @@ -333,5 +359,86 @@ def reducer(tensor): assert audit_metrics["rlk_audit_worst_rollout_id"].item() == pytest.approx(7.0) +@pytest.mark.unit +def test_policy_loss_uses_consistency_metadata_for_audit_context(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() + + args = _policy_args("audit") + args.model_name = None + args.train_backend = None + args.rlk_contract_id = None + args.rlk_batch_layout_fingerprint = None + args.rlk_provenance_fingerprint = None + batch = _policy_batch() + batch["consistency_metadata"] = [_complete_consistency_record()] + batch["consistency_batch_layout_fingerprints"] = [{"fingerprint": "record-layout"}] + + _, metrics = megatron_loss_module.policy_loss_function( + args, + batch, + torch.zeros(1, 2, 4), + reducer, + ) + + assert metrics["rlk_audit_warning_count"].item() == pytest.approx(0.0) + assert metrics["rlk_audit_metadata_warning_count"].item() == pytest.approx(0.0) + assert metrics["rlk_audit_replay_case_count"].item() == pytest.approx(5.0) + assert metrics["rlk_audit_worst_sample_index"].item() == pytest.approx(42.0) + + +@pytest.mark.unit +def test_policy_loss_adds_linear_logp_runtime_provenance(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) + + runtime_provenance = { + "operator": "linear_logp", + "requested_backend": "registry", + "actual_backend": "vime.native.linear_logp", + "fallback": True, + "fallback_reason": "unit fallback", + } + + 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) + monkeypatch.setattr(megatron_loss_module, "get_linear_logp_runtime_metadata", lambda: runtime_provenance) + + def reducer(tensor): + return tensor.mean() + + batch = _policy_batch() + batch["consistency_metadata"] = [_complete_consistency_record()] + batch["consistency_batch_layout_fingerprints"] = [{"fingerprint": "record-layout"}] + _, metrics = megatron_loss_module.policy_loss_function( + _policy_args("audit"), + batch, + torch.zeros(1, 2, 4), + reducer, + rl_kernel_linear_logp_context=object(), + ) + + assert metrics["rlk_audit_runtime_fallback"].item() == pytest.approx(1.0) + assert metrics["rlk_audit_metadata_warning_count"].item() == pytest.approx(1.0) + + if __name__ == "__main__": raise SystemExit(pytest.main([__file__])) diff --git a/vime/backends/megatron_utils/loss.py b/vime/backends/megatron_utils/loss.py index ed3762d7..72b7554d 100644 --- a/vime/backends/megatron_utils/loss.py +++ b/vime/backends/megatron_utils/loss.py @@ -9,8 +9,9 @@ from megatron.core import mpu from torch.utils.checkpoint import checkpoint +from vime.utils.consistency_audit import run_consistency_audit 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.dlogp_diagnostics import is_dlogp_audit_enabled from vime.utils.misc import load_function from vime.utils.ppo_utils import ( calculate_log_probs_and_entropy, @@ -35,6 +36,7 @@ ) from .rl_kernel import ( LinearLogpContext, + get_linear_logp_runtime_metadata, get_rl_kernel_fallback_count, maybe_compute_linear_logp, warn_linear_logp_fallback, @@ -1181,19 +1183,23 @@ def policy_loss_function( dlogp_audit_metrics = {} if is_dlogp_audit_enabled(args): - dlogp_audit_metrics = compute_dlogp_diagnostics( + dlogp_audit_metrics = run_consistency_audit( audit_train_log_probs, batch.get("rollout_log_probs"), batch["loss_masks"], + args=args, + batch=batch, 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), + runtime_provenance=( + get_linear_logp_runtime_metadata() if rl_kernel_linear_logp_context is not None else None + ), eps_clip=getattr(args, "eps_clip", 0.2), ).metrics diff --git a/vime/ray/rollout.py b/vime/ray/rollout.py index 972a3ad8..ba139b92 100644 --- a/vime/ray/rollout.py +++ b/vime/ray/rollout.py @@ -23,6 +23,7 @@ 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_audit import build_consistency_replay_manifest from vime.utils.consistency_metadata import ( build_batch_layout_fingerprints, get_consistency_mode, @@ -866,7 +867,8 @@ 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: + consistency_mode = get_consistency_mode(self.args) + if consistency_mode != "off" or "consistency_metadata" in data: data["consistency_batch_layout_fingerprints"] = build_batch_layout_fingerprints( data, partitions=partitions, @@ -915,6 +917,13 @@ def _split_train_data_by_dp(self, data): rollout_data["global_batch_sizes"] = global_batch_sizes rollout_data["num_microbatches"] = num_microbatches rollout_data["micro_batch_indices"] = micro_batch_indices[r] + if consistency_mode != "off" or "consistency_metadata" in rollout_data: + rollout_data["consistency_replay_manifest"] = build_consistency_replay_manifest( + rollout_data, + mode=consistency_mode, + rank=r, + validation=data.get("consistency_metadata_validation"), + ) _tensorize_rollout_data_for_training(rollout_data) transport = getattr(self.args, "rollout_data_transport", "object-store") if transport == "nixl": diff --git a/vime/utils/consistency_audit.py b/vime/utils/consistency_audit.py new file mode 100644 index 00000000..c7dec738 --- /dev/null +++ b/vime/utils/consistency_audit.py @@ -0,0 +1,836 @@ +"""Reusable rollout-to-training consistency audit harness. + +The Phase 1 modules own compact metadata and raw diagnostic math. This module +ties those pieces together for Phase 3: validate comparison preconditions, +compute read-only dlogp diagnostics, and emit lightweight replay/result-cube +records that existing debug dumps can carry. +""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from typing import Any + +import torch + +from vime.utils.consistency_metadata import ( + CONSISTENCY_METADATA_SCHEMA_VERSION, + ConsistencyMetadataIssue, + ConsistencyMetadataValidation, + raise_for_consistency_metadata_failures, + stable_fingerprint, + validate_samples_consistency_metadata, +) +from vime.utils.dlogp_diagnostics import ( + DlogpAuditReport, + compute_dlogp_diagnostics, + get_rlk_consistency_mode, +) + +CONSISTENCY_AUDIT_SCHEMA_VERSION = 1 +AUDIT_REQUIRED_METADATA_FIELDS = ( + ("sample.session_id", "session_id_missing"), + ("batch_layout.fingerprint", "batch_layout_missing"), + ("position_cache.fingerprint", "position_cache_metadata_missing"), + ("quantization.fingerprint", "quantization_metadata_missing"), + ("weight.pre_update", "weight_update_status_missing"), +) + + +@dataclass(frozen=True) +class ConsistencyAuditResult: + metrics: dict[str, torch.Tensor] + dlogp_report: DlogpAuditReport + metadata_validation: ConsistencyMetadataValidation + replay_manifest: dict[str, Any] + result_cube: dict[str, Any] + diagnostic_metadata: tuple[dict[str, Any] | None, ...] = field(default_factory=tuple) + + +@dataclass +class _BatchMetadataSample: + consistency_metadata: dict[str, Any] | None + index: int | None = None + rollout_id: int | None = None + + +def run_consistency_audit( + train_log_probs: Sequence[torch.Tensor] | torch.Tensor, + rollout_log_probs: Sequence[torch.Tensor] | torch.Tensor | None, + loss_masks: Sequence[torch.Tensor] | torch.Tensor, + *, + args: Any | None = None, + batch: Mapping[str, Any] | None = None, + sample_indices: Sequence[int | torch.Tensor] | torch.Tensor | None = None, + rollout_ids: Sequence[int | torch.Tensor] | torch.Tensor | 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, + runtime_provenance: Mapping[str, Any] | None = None, + eps_clip: float = 0.2, + prefix: str = "rlk_audit_", +) -> ConsistencyAuditResult: + """Run the read-only consistency audit over already teacher-forced logprobs.""" + + mode = get_rlk_consistency_mode(args) + runtime_provenance = _plain_mapping(runtime_provenance) + validation = validate_consistency_audit_batch( + batch, + mode=mode, + runtime_provenance=runtime_provenance, + ) + if mode == "strict": + raise_for_consistency_metadata_failures(validation) + + diagnostic_metadata = build_consistency_diagnostic_metadata(batch) + report = compute_dlogp_diagnostics( + train_log_probs, + rollout_log_probs, + loss_masks, + sample_indices=sample_indices if sample_indices is not None else _batch_get(batch, "sample_indices"), + rollout_ids=rollout_ids if rollout_ids is not None else _batch_get(batch, "rollout_ids"), + metadata=diagnostic_metadata or None, + rank=rank, + model_name=model_name if model_name is not None else getattr(args, "model_name", None), + backend_id=backend_id if backend_id is not None else getattr(args, "train_backend", "megatron"), + contract_id=contract_id if contract_id is not None else getattr(args, "rlk_contract_id", None), + batch_layout_fingerprint=( + batch_layout_fingerprint + if batch_layout_fingerprint is not None + else getattr(args, "rlk_batch_layout_fingerprint", None) + ), + provenance_fingerprint=( + provenance_fingerprint + if provenance_fingerprint is not None + else _first_present( + getattr(args, "rlk_provenance_fingerprint", None), + stable_fingerprint(runtime_provenance) if runtime_provenance else None, + ) + ), + eps_clip=eps_clip, + prefix=prefix, + ) + replay_manifest = build_consistency_replay_manifest( + batch, + mode=mode, + rank=rank, + validation=validation.to_dict(), + diagnostic_metadata=diagnostic_metadata, + runtime_provenance=runtime_provenance, + ) + result_cube = build_consistency_result_cube( + args=args, + batch=batch, + mode=mode, + rank=rank, + dlogp_report=report, + validation=validation, + diagnostic_metadata=diagnostic_metadata, + runtime_provenance=runtime_provenance, + ) + + metrics = dict(report.metrics) + metrics.update( + _audit_bookkeeping_metrics( + report, + validation=validation, + replay_manifest=replay_manifest, + result_cube=result_cube, + runtime_provenance=runtime_provenance, + prefix=prefix, + ) + ) + return ConsistencyAuditResult( + metrics={key: value.clone().detach() for key, value in metrics.items()}, + dlogp_report=report, + metadata_validation=validation, + replay_manifest=replay_manifest, + result_cube=result_cube, + diagnostic_metadata=tuple(diagnostic_metadata), + ) + + +def validate_consistency_audit_batch( + batch: Mapping[str, Any] | None, + *, + mode: str, + runtime_provenance: Mapping[str, Any] | None = None, +) -> ConsistencyMetadataValidation: + """Validate batch-level consistency metadata before drift attribution.""" + + mode = str(mode).lower() + if mode == "off": + return ConsistencyMetadataValidation(mode=mode) + + if batch is None: + return _single_issue_validation( + mode, + code="consistency_batch_missing", + message="Training batch is unavailable for consistency metadata validation.", + ) + + precomputed = batch.get("consistency_metadata_validation") + if isinstance(precomputed, Mapping): + validation = _metadata_validation_from_dict(precomputed, mode=mode) + validation = _with_audit_required_metadata_issues(batch, validation, mode=mode) + return _with_runtime_provenance_issues(validation, runtime_provenance, mode=mode) + + records = _as_optional_mapping_list(batch.get("consistency_metadata")) + if records: + samples = [ + _BatchMetadataSample( + consistency_metadata=record, + index=_sequence_value(batch.get("sample_indices"), i), + rollout_id=_sequence_value(batch.get("rollout_ids"), i), + ) + for i, record in enumerate(records) + ] + validation = validate_samples_consistency_metadata(samples, mode=mode) + validation = _with_audit_required_metadata_issues(batch, validation, mode=mode) + return _with_runtime_provenance_issues(validation, runtime_provenance, mode=mode) + + validation = _single_issue_validation( + mode, + code="consistency_metadata_missing", + message="Training batch is missing consistency metadata required for audit/strict comparison.", + ) + return _with_runtime_provenance_issues(validation, runtime_provenance, mode=mode) + + +def build_consistency_diagnostic_metadata(batch: Mapping[str, Any] | None) -> tuple[dict[str, Any] | None, ...]: + """Flatten Phase 1 metadata into the context fields dlogp diagnostics need.""" + + if batch is None: + return () + + records = _as_optional_mapping_list(batch.get("consistency_metadata")) + layouts = _as_optional_mapping_list(batch.get("consistency_batch_layout_fingerprints")) + sample_count = max(len(records), len(layouts), _batch_sample_count(batch)) + if sample_count == 0: + return () + + result: list[dict[str, Any] | None] = [] + for i in range(sample_count): + record = records[i] if i < len(records) else None + layout = layouts[i] if i < len(layouts) else None + if record is None and layout is None: + result.append(None) + continue + + provenance = _mapping_at(record, "provenance") + actual = _mapping_at(provenance, "actual") + old_logp = _mapping_at(record, "old_logp") + model = _mapping_at(record, "model") + flattened = { + "model_name": model.get("name"), + "backend_id": _first_present( + actual.get("backend"), + actual.get("actual_backend"), + actual.get("backend_id"), + old_logp.get("source"), + ), + "contract_id": old_logp.get("contract_id"), + "batch_layout_fingerprint": _first_present( + None if layout is None else layout.get("fingerprint"), + _path(record, "batch_layout.fingerprint"), + ), + "provenance_fingerprint": _first_present( + provenance.get("actual_fingerprint"), + provenance.get("requested_fingerprint"), + ), + "router_policy": actual.get("router_policy"), + "vllm_enable_prefix_caching": actual.get("vllm_enable_prefix_caching"), + "vllm_enable_deterministic_inference": actual.get("vllm_enable_deterministic_inference"), + "tensor_model_parallel_size": actual.get("tensor_model_parallel_size"), + "megatron_tensor_parallel_size": actual.get("megatron_tensor_parallel_size"), + "context_parallel_size": actual.get("context_parallel_size"), + "megatron_context_parallel_size": actual.get("megatron_context_parallel_size"), + "consistency_metadata_fingerprint": None if record is None else record.get("fingerprint"), + "dynamic_sampling": None if record is None else record.get("dynamic_sampling"), + } + result.append(flattened) + return tuple(result) + + +def build_consistency_replay_manifest( + batch: Mapping[str, Any] | None, + *, + mode: str, + rank: int | None = None, + validation: Mapping[str, Any] | None = None, + diagnostic_metadata: Sequence[Mapping[str, Any] | None] | None = None, + runtime_provenance: Mapping[str, Any] | None = None, + max_samples: int | None = None, +) -> dict[str, Any]: + """Build a lightweight replay/debug manifest for existing train-data dumps.""" + + sample_count = _batch_sample_count(batch) + if max_samples is not None: + sample_count = min(sample_count, max(0, int(max_samples))) + diagnostic_metadata = tuple(diagnostic_metadata or build_consistency_diagnostic_metadata(batch)) + runtime_provenance = _plain_mapping(runtime_provenance) + + samples = [] + for position in range(sample_count): + record = _record_at(batch, "consistency_metadata", position) + layout = _record_at(batch, "consistency_batch_layout_fingerprints", position) + response_length = _int_or_none(_sequence_value(_batch_get(batch, "response_lengths"), position)) + loss_mask = _batch_sequence_value(batch, "loss_masks", position) + diag = diagnostic_metadata[position] if position < len(diagnostic_metadata) else None + samples.append( + { + "sample_position": position, + "sample_index": _sequence_value(_batch_get(batch, "sample_indices"), position), + "rollout_id": _sequence_value(_batch_get(batch, "rollout_ids"), position), + "total_length": _int_or_none(_sequence_value(_batch_get(batch, "total_lengths"), position)), + "response_length": response_length, + "active_token_count": _active_token_count(loss_mask, response_length), + "has_rollout_log_probs": _batch_sequence_value(batch, "rollout_log_probs", position) + is not None, + "consistency_metadata_fingerprint": None if record is None else record.get("fingerprint"), + "batch_layout_fingerprint": _first_present( + None if layout is None else layout.get("fingerprint"), + None if diag is None else diag.get("batch_layout_fingerprint"), + ), + "provenance_fingerprint": None if diag is None else diag.get("provenance_fingerprint"), + "dynamic_sampling": _path(record, "dynamic_sampling"), + } + ) + + manifest = { + "schema_version": CONSISTENCY_AUDIT_SCHEMA_VERSION, + "metadata_schema_version": CONSISTENCY_METADATA_SCHEMA_VERSION, + "mode": mode, + "rank": rank, + "sample_count": len(samples), + "samples": samples, + "batch_invariance_cases": build_batch_invariance_replay_cases( + batch, + diagnostic_metadata=diagnostic_metadata, + max_samples=max_samples, + ), + "validation": dict(validation or {}), + "runtime_provenance": runtime_provenance, + "runtime_provenance_fingerprint": stable_fingerprint(runtime_provenance) if runtime_provenance else None, + } + manifest["fingerprint"] = stable_fingerprint(manifest) + return manifest + + +def build_batch_invariance_replay_cases( + batch: Mapping[str, Any] | None, + *, + diagnostic_metadata: Sequence[Mapping[str, Any] | None] | None = None, + max_samples: int | None = None, +) -> list[dict[str, Any]]: + """Describe replay cases that keep one sample fixed while varying layout.""" + + sample_count = _batch_sample_count(batch) + if max_samples is not None: + sample_count = min(sample_count, max(0, int(max_samples))) + diagnostic_metadata = tuple(diagnostic_metadata or build_consistency_diagnostic_metadata(batch)) + + cases: list[dict[str, Any]] = [] + for position in range(sample_count): + sample_ref = { + "sample_position": position, + "sample_index": _sequence_value(_batch_get(batch, "sample_indices"), position), + "rollout_id": _sequence_value(_batch_get(batch, "rollout_ids"), position), + } + layout = _record_at(batch, "consistency_batch_layout_fingerprints", position) or {} + record = _record_at(batch, "consistency_metadata", position) or {} + diag = diagnostic_metadata[position] if position < len(diagnostic_metadata) else None + base = { + **sample_ref, + "batch_layout_fingerprint": _first_present( + layout.get("fingerprint"), + None if diag is None else diag.get("batch_layout_fingerprint"), + ), + "consistency_metadata_fingerprint": record.get("fingerprint"), + } + cases.extend( + [ + { + **base, + "case": "same_sample_alone", + "varied_axes": ("batch_size", "neighboring_samples"), + "expected": "same-sample dlogp stays within the declared tolerance", + }, + { + **base, + "case": "same_sample_mixed_batch", + "varied_axes": ("batch_order", "microbatch_membership"), + "expected": "mixed-batch placement does not change the fixed sample", + }, + { + **base, + "case": "padding_packing_variant", + "varied_axes": ("padding_side", "packed_order", "microbatch_offset"), + "expected": "padding and packing changes do not alter active-token logprobs", + }, + { + **base, + "case": "active_token_density_variant", + "varied_axes": ("active_mask_density",), + "active_mask_density": layout.get("active_mask_density"), + "expected": "neighboring active-token density does not alter the fixed sample", + }, + { + **base, + "case": "dynamic_sampling_keep_drop_variant", + "varied_axes": ("dynamic_sampling_keep_drop",), + "dynamic_sampling": record.get("dynamic_sampling"), + "expected": "keep/drop decisions do not hide same-sample drift", + }, + ] + ) + return cases + + +def build_consistency_result_cube( + *, + args: Any | None, + batch: Mapping[str, Any] | None, + mode: str, + rank: int | None, + dlogp_report: DlogpAuditReport, + validation: ConsistencyMetadataValidation, + diagnostic_metadata: Sequence[Mapping[str, Any] | None], + runtime_provenance: Mapping[str, Any] | None = None, +) -> dict[str, Any]: + """Return a compact result-cube entry indexed by normalized audit axes.""" + + runtime_provenance = _plain_mapping(runtime_provenance) + axes = { + "batch_layout": _unique_axis(diagnostic_metadata, "batch_layout_fingerprint"), + "dtype": _dtype_axis(args), + "router_policy": _axis_from_args_or_metadata(args, diagnostic_metadata, "router_policy"), + "cache_policy": _first_present( + getattr(args, "vllm_enable_prefix_caching", None), + _axis_from_args_or_metadata(args, diagnostic_metadata, "vllm_enable_prefix_caching"), + ), + "tp": _first_present( + getattr(args, "tensor_model_parallel_size", None), + getattr(args, "megatron_tensor_parallel_size", None), + _unique_axis(diagnostic_metadata, "tensor_model_parallel_size"), + _unique_axis(diagnostic_metadata, "megatron_tensor_parallel_size"), + ), + "sp": _first_present(getattr(args, "sequence_parallel", None), getattr(args, "use_sequence_parallel", None)), + "cp": _first_present( + getattr(args, "context_parallel_size", None), + getattr(args, "megatron_context_parallel_size", None), + _unique_axis(diagnostic_metadata, "context_parallel_size"), + _unique_axis(diagnostic_metadata, "megatron_context_parallel_size"), + ), + "logp_backend": _first_present( + runtime_provenance.get("backend_id"), + runtime_provenance.get("actual_backend"), + _unique_axis(diagnostic_metadata, "backend_id"), + ), + "deterministic_policy": _first_present( + getattr(args, "vllm_enable_deterministic_inference", None), + getattr(args, "rlk_deterministic_logp", None), + ), + } + metric_summary = { + "active_token_count": _metric_float(dlogp_report, "rlk_audit_active_token_count"), + "max_abs_dlogp": _metric_float(dlogp_report, "rlk_audit_dlogp_abs_max"), + "warning_count": _metric_float(dlogp_report, "rlk_audit_warning_count"), + "metadata_warning_count": float(len(validation.warnings)), + "metadata_failure_count": float(len(validation.failures)), + "sample_count": float(_batch_sample_count(batch)), + "runtime_fallback": 1.0 if runtime_provenance.get("fallback") else 0.0, + "runtime_strict_failure": 1.0 if runtime_provenance.get("strict_failure") else 0.0, + } + entry = { + "schema_version": CONSISTENCY_AUDIT_SCHEMA_VERSION, + "mode": mode, + "rank": rank, + "axes": axes, + "metrics": metric_summary, + "worst_token": dlogp_report.worst_token, + "metadata_validation": validation.to_dict(), + "runtime_provenance": runtime_provenance, + "runtime_provenance_fingerprint": stable_fingerprint(runtime_provenance) if runtime_provenance else None, + } + entry["fingerprint"] = stable_fingerprint(entry) + return entry + + +def _audit_bookkeeping_metrics( + report: DlogpAuditReport, + *, + validation: ConsistencyMetadataValidation, + replay_manifest: Mapping[str, Any], + result_cube: Mapping[str, Any], + runtime_provenance: Mapping[str, Any], + prefix: str, +) -> dict[str, torch.Tensor]: + device, dtype = _metric_device_dtype(report) + return { + f"{prefix}metadata_warning_count": torch.tensor(float(len(validation.warnings)), device=device, dtype=dtype), + f"{prefix}metadata_failure_count": torch.tensor(float(len(validation.failures)), device=device, dtype=dtype), + f"{prefix}metadata_active_token_count": torch.tensor( + float(validation.active_token_count), + device=device, + dtype=dtype, + ), + f"{prefix}replay_case_count": torch.tensor( + float(len(replay_manifest.get("batch_invariance_cases", ()))), + device=device, + dtype=dtype, + ), + f"{prefix}result_cube_axis_count": torch.tensor( + float(len(result_cube.get("axes", {}))), + device=device, + dtype=dtype, + ), + f"{prefix}runtime_fallback": torch.tensor( + 1.0 if runtime_provenance.get("fallback") else 0.0, + device=device, + dtype=dtype, + ), + f"{prefix}runtime_strict_failure": torch.tensor( + 1.0 if runtime_provenance.get("strict_failure") else 0.0, + device=device, + dtype=dtype, + ), + } + + +def _metadata_validation_from_dict(record: Mapping[str, Any], *, mode: str) -> ConsistencyMetadataValidation: + warnings = [_issue_from_dict(issue, default_severity="warning") for issue in record.get("warnings", ())] + failures = [_issue_from_dict(issue, default_severity="error") for issue in record.get("failures", ())] + return ConsistencyMetadataValidation( + mode=mode, + active_token_count=int(record.get("active_token_count") or 0), + zero_active_token_samples=list(record.get("zero_active_token_samples") or ()), + warnings=warnings, + failures=failures, + ) + + +def _issue_from_dict(record: Mapping[str, Any], *, default_severity: str) -> ConsistencyMetadataIssue: + return ConsistencyMetadataIssue( + code=str(record.get("code") or "consistency_metadata_issue"), + message=str(record.get("message") or ""), + severity=str(record.get("severity") or default_severity), + sample_index=_int_or_none(record.get("sample_index")), + rollout_id=_int_or_none(record.get("rollout_id")), + field=record.get("field"), + ) + + +def _single_issue_validation(mode: str, *, code: str, message: str) -> ConsistencyMetadataValidation: + severity = "error" if mode == "strict" else "warning" + issue = ConsistencyMetadataIssue(code=code, message=message, severity=severity) + if mode == "strict": + return ConsistencyMetadataValidation(mode=mode, failures=[issue]) + return ConsistencyMetadataValidation(mode=mode, warnings=[issue]) + + +def _with_audit_required_metadata_issues( + batch: Mapping[str, Any], + validation: ConsistencyMetadataValidation, + *, + mode: str, +) -> ConsistencyMetadataValidation: + records = _as_optional_mapping_list(batch.get("consistency_metadata")) + if not records: + return validation + + warnings = list(validation.warnings) + failures = list(validation.failures) + seen = { + (issue.code, issue.sample_index, issue.rollout_id, issue.field) + for issue in (*warnings, *failures) + } + + for position, record in enumerate(records): + if record is None: + continue + layout = _record_at(batch, "consistency_batch_layout_fingerprints", position) or {} + sample_index = _first_present( + _path(record, "sample.index"), + _sequence_value(batch.get("sample_indices"), position), + ) + rollout_id = _first_present( + _path(record, "sample.rollout_id"), + _sequence_value(batch.get("rollout_ids"), position), + ) + for field_path, code in AUDIT_REQUIRED_METADATA_FIELDS: + value = _first_present( + layout.get("fingerprint") if field_path == "batch_layout.fingerprint" else None, + _path(record, field_path), + ) + if value not in (None, ""): + continue + key = (code, _int_or_none(sample_index), _int_or_none(rollout_id), field_path) + if key in seen: + continue + issue = ConsistencyMetadataIssue( + code=code, + message=f"Consistency audit metadata field {field_path!r} is missing.", + severity="error" if mode == "strict" else "warning", + sample_index=_int_or_none(sample_index), + rollout_id=_int_or_none(rollout_id), + field=field_path, + ) + if mode == "strict": + failures.append(issue) + else: + warnings.append(issue) + seen.add(key) + + return ConsistencyMetadataValidation( + mode=validation.mode, + active_token_count=validation.active_token_count, + zero_active_token_samples=list(validation.zero_active_token_samples), + warnings=warnings, + failures=failures, + ) + + +def _with_runtime_provenance_issues( + validation: ConsistencyMetadataValidation, + runtime_provenance: Mapping[str, Any] | None, + *, + mode: str, +) -> ConsistencyMetadataValidation: + if mode == "off" or not runtime_provenance: + return validation + + warnings = list(validation.warnings) + failures = list(validation.failures) + issues: list[ConsistencyMetadataIssue] = [] + if runtime_provenance.get("fallback"): + issues.append( + ConsistencyMetadataIssue( + code="undeclared_linear_logp_runtime_fallback", + message="Training linear_logp runtime reported fallback during consistency audit.", + severity="error" if mode == "strict" else "warning", + field="linear_logp_runtime.fallback", + ) + ) + if runtime_provenance.get("strict_failure"): + issues.append( + ConsistencyMetadataIssue( + code="linear_logp_runtime_strict_failure", + message="Training linear_logp runtime reported a strict failure.", + severity="error", + field="linear_logp_runtime.strict_failure", + ) + ) + + for issue in issues: + if issue.severity == "error": + failures.append(issue) + else: + warnings.append(issue) + + return ConsistencyMetadataValidation( + mode=validation.mode, + active_token_count=validation.active_token_count, + zero_active_token_samples=list(validation.zero_active_token_samples), + warnings=warnings, + failures=failures, + ) + + +def _metric_device_dtype(report: DlogpAuditReport) -> tuple[torch.device, torch.dtype]: + for metric in report.metrics.values(): + return metric.device, metric.dtype + return torch.device("cpu"), torch.float32 + + +def _metric_float(report: DlogpAuditReport, key: str) -> float: + value = report.metrics.get(key) + if value is None: + return 0.0 + return float(value.detach().cpu().item()) + + +def _as_optional_mapping_list(value: Any) -> tuple[dict[str, Any] | None, ...]: + if value is None: + return () + if isinstance(value, Mapping): + return (dict(value),) + result = [] + for item in value: + if item is None: + result.append(None) + elif isinstance(item, Mapping): + result.append(dict(item)) + else: + result.append(None) + return tuple(result) + + +def _record_at(batch: Mapping[str, Any] | None, key: str, position: int) -> dict[str, Any] | None: + records = _as_optional_mapping_list(_batch_get(batch, key)) + if position >= len(records): + return None + return records[position] + + +def _batch_get(batch: Mapping[str, Any] | None, key: str) -> Any: + if batch is None: + return None + return batch.get(key) + + +def _batch_sample_count(batch: Mapping[str, Any] | None) -> int: + if batch is None: + return 0 + for key in ( + "tokens", + "unconcat_tokens", + "response_lengths", + "loss_masks", + "rollout_log_probs", + "sample_indices", + "rollout_ids", + "consistency_metadata", + ): + value = batch.get(key) + if value is not None: + return _sample_count_from_value(value, key) + return 0 + + +def _sample_count_from_value(value: Any, key: str) -> int: + if isinstance(value, torch.Tensor): + if key in {"tokens", "unconcat_tokens", "loss_masks", "rollout_log_probs"}: + return 1 if value.ndim <= 1 else int(value.shape[0]) + return int(value.numel()) if value.ndim <= 1 else int(value.shape[0]) + try: + return len(value) + except TypeError: + return 1 + + +def _batch_sequence_value(batch: Mapping[str, Any] | None, key: str, position: int) -> Any: + value = _batch_get(batch, key) + if isinstance(value, torch.Tensor) and key in {"tokens", "unconcat_tokens", "loss_masks", "rollout_log_probs"}: + if value.ndim == 0: + return value.detach().cpu().item() if position == 0 else None + if value.ndim == 1: + return value if position == 0 else None + if position >= value.shape[0]: + return None + return value[position] + return _sequence_value(value, position) + + +def _sequence_value(value: Any, position: int) -> Any: + if value is None: + return None + if isinstance(value, torch.Tensor): + flat = value.detach().cpu().flatten() + if position >= flat.numel(): + return None + return flat[position].item() + try: + if position >= len(value): + return None + item = value[position] + except TypeError: + return value if position == 0 else None + if isinstance(item, torch.Tensor): + if item.numel() == 1: + return item.detach().cpu().item() + return item.detach().cpu().tolist() + return item + + +def _active_token_count(loss_mask: Any, response_length: int | None) -> int | None: + if loss_mask is None: + return response_length + if isinstance(loss_mask, torch.Tensor): + return int(loss_mask.detach().cpu().to(dtype=torch.bool).sum().item()) + try: + return sum(1 for item in loss_mask if item) + except TypeError: + return None + + +def _mapping_at(value: Mapping[str, Any] | None, key: str) -> dict[str, Any]: + if not isinstance(value, Mapping): + return {} + item = value.get(key) + return dict(item) if isinstance(item, Mapping) else {} + + +def _path(value: Mapping[str, Any] | None, path: str) -> Any: + current: Any = value + for part in path.split("."): + if not isinstance(current, Mapping): + return None + current = current.get(part) + return current + + +def _first_present(*values: Any) -> Any: + for value in values: + if value is not None and value != "": + return value + return None + + +def _unique_axis(records: Sequence[Mapping[str, Any] | None], field: str) -> Any: + values = [] + for record in records: + if record is not None and record.get(field) is not None: + values.append(record[field]) + unique = list(dict.fromkeys(values)) + if not unique: + return None + if len(unique) == 1: + return unique[0] + return unique + + +def _axis_from_args_or_metadata(args: Any | None, records: Sequence[Mapping[str, Any] | None], field: str) -> Any: + value = getattr(args, field, None) if args is not None else None + if value is not None: + return value + return _unique_axis(records, field) + + +def _dtype_axis(args: Any | None) -> Any: + if args is None: + return None + for attr in ("params_dtype", "dtype"): + value = getattr(args, attr, None) + if value is not None: + return _normalize_dtype_value(value) + if getattr(args, "bf16", False): + return "bf16" + if getattr(args, "fp16", False): + return "fp16" + return None + + +def _normalize_dtype_value(value: Any) -> Any: + if isinstance(value, torch.dtype): + return str(value).replace("torch.", "") + if isinstance(value, str): + return value.replace("torch.", "") + return value + + +def _plain_mapping(value: Mapping[str, Any] | None) -> dict[str, Any]: + return dict(value or {}) + + +def _int_or_none(value: Any) -> int | None: + if value is None or value == "": + return None + try: + return int(value) + except (TypeError, ValueError): + return None